2024-07-17 08:15:20 +08:00
import os
import sys
2023-11-24 12:27:44 +08:00
import traceback
2024-05-03 03:24:49 +08:00
from unittest import mock
2024-07-17 08:15:20 +08:00
2023-11-24 12:27:44 +08:00
from dotenv import load_dotenv
2024-07-17 08:15:20 +08:00
import litellm . proxy
import litellm . proxy . proxy_server
2023-11-24 12:27:44 +08:00
load_dotenv ( )
2024-07-17 08:15:20 +08:00
import io
import os
2023-11-24 12:27:44 +08:00
# this file is to test litellm/proxy
sys . path . insert (
0 , os . path . abspath ( " ../.. " )
2023-12-25 16:40:38 +08:00
) # Adds the parent directory to the system path
2024-07-17 08:15:20 +08:00
import asyncio
import logging
import pytest
2023-11-24 12:27:44 +08:00
import litellm
2024-07-17 08:15:20 +08:00
from litellm import RateLimitError , Timeout , completion , completion_cost , embedding
2023-12-25 16:40:38 +08:00
2023-12-07 09:40:38 +08:00
# Configure logging
logging . basicConfig (
level = logging . DEBUG , # Set the desired logging level
format = " %(asctime)s - %(levelname)s - %(message)s " ,
)
2023-11-24 12:27:44 +08:00
2024-08-08 06:37:02 +08:00
from unittest . mock import AsyncMock , patch
2024-07-17 08:15:20 +08:00
from fastapi import FastAPI
2023-11-24 12:27:44 +08:00
# test /chat/completion request to the proxy
from fastapi . testclient import TestClient
2024-07-17 08:15:20 +08:00
from litellm . integrations . custom_logger import CustomLogger
from litellm . proxy . proxy_server import ( # Replace with the actual module where your FastAPI router is defined
2024-06-16 06:09:49 +08:00
app ,
2023-12-25 16:40:38 +08:00
initialize ,
2024-07-17 08:15:20 +08:00
save_worker_config ,
)
from litellm . proxy . utils import ProxyLogging
2023-12-06 03:13:09 +08:00
2023-12-07 10:38:44 +08:00
# Your bearer token
2024-01-27 12:15:35 +08:00
token = " sk-1234 "
2023-12-07 10:38:44 +08:00
2023-12-25 16:40:38 +08:00
headers = { " Authorization " : f " Bearer { token } " }
2024-05-03 03:47:27 +08:00
example_completion_result = {
" choices " : [
{
" message " : {
" content " : " Whispers of the wind carry dreams to me. " ,
2024-06-02 18:49:34 +08:00
" role " : " assistant " ,
2024-05-03 03:47:27 +08:00
}
}
] ,
}
2024-05-03 04:36:23 +08:00
example_embedding_result = {
2024-06-02 18:49:34 +08:00
" object " : " list " ,
" data " : [
{
" object " : " embedding " ,
" index " : 0 ,
" embedding " : [
- 0.006929283495992422 ,
- 0.005336422007530928 ,
- 4.547132266452536e-05 ,
- 0.024047505110502243 ,
- 0.006929283495992422 ,
- 0.005336422007530928 ,
- 4.547132266452536e-05 ,
- 0.024047505110502243 ,
- 0.006929283495992422 ,
- 0.005336422007530928 ,
- 4.547132266452536e-05 ,
- 0.024047505110502243 ,
] ,
}
] ,
" model " : " text-embedding-3-small " ,
" usage " : { " prompt_tokens " : 5 , " total_tokens " : 5 } ,
2024-05-03 04:36:23 +08:00
}
example_image_generation_result = {
2024-06-02 18:49:34 +08:00
" created " : 1589478378 ,
" data " : [ { " url " : " https://... " } , { " url " : " https://... " } ] ,
2024-05-03 04:36:23 +08:00
}
2024-05-03 03:47:27 +08:00
2023-12-25 16:40:38 +08:00
2024-05-03 03:24:49 +08:00
def mock_patch_acompletion ( ) :
2024-05-03 04:36:23 +08:00
return mock . patch (
2024-05-03 03:24:49 +08:00
" litellm.proxy.proxy_server.llm_router.acompletion " ,
2024-05-03 03:47:27 +08:00
return_value = example_completion_result ,
2024-05-03 04:36:23 +08:00
)
def mock_patch_aembedding ( ) :
return mock . patch (
" litellm.proxy.proxy_server.llm_router.aembedding " ,
return_value = example_embedding_result ,
)
def mock_patch_aimage_generation ( ) :
return mock . patch (
" litellm.proxy.proxy_server.llm_router.aimage_generation " ,
return_value = example_image_generation_result ,
)
2024-05-03 03:24:49 +08:00
2023-12-12 13:30:02 +08:00
@pytest.fixture ( scope = " function " )
Set fake env vars for `client_no_auth` fixture
This allows all of the tests in `test_proxy_server.py` to pass, with the
exception of `test_load_router_config`, without needing to set up real
environment variables.
Before:
```shell
$ env -i PATH=$PATH poetry run pytest litellm/tests/test_proxy_server.py -k 'not test_load_router_config' --disable-warnings
...
========================================================== short test summary info ===========================================================
ERROR litellm/tests/test_proxy_server.py::test_bedrock_embedding - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_chat_completion - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_chat_completion_azure - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_chat_completion_optional_params - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_embedding - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_engines_model_chat_completions - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_health - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_img_gen - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_openai_deployments_model_chat_completions_azure - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
========================================== 2 skipped, 1 deselected, 39 warnings, 9 errors in 3.24s ===========================================
```
After:
```shell
$ env -i PATH=$PATH poetry run pytest litellm/tests/test_proxy_server.py -k 'not test_load_router_config' --disable-warnings
============================================================ test session starts =============================================================
platform darwin -- Python 3.12.3, pytest-7.4.4, pluggy-1.5.0
rootdir: /Users/abramowi/Code/OpenSource/litellm
plugins: anyio-4.3.0, asyncio-0.23.6, mock-3.14.0
asyncio: mode=Mode.STRICT
collected 12 items / 1 deselected / 11 selected
litellm/tests/test_proxy_server.py s.........s [100%]
========================================== 9 passed, 2 skipped, 1 deselected, 48 warnings in 8.42s ===========================================
```
2024-05-12 06:22:30 +08:00
def fake_env_vars ( monkeypatch ) :
# Set some fake environment variables
monkeypatch . setenv ( " OPENAI_API_KEY " , " fake_openai_api_key " )
monkeypatch . setenv ( " OPENAI_API_BASE " , " http://fake-openai-api-base " )
monkeypatch . setenv ( " AZURE_API_BASE " , " http://fake-azure-api-base " )
monkeypatch . setenv ( " AZURE_OPENAI_API_KEY " , " fake_azure_openai_api_key " )
monkeypatch . setenv ( " AZURE_SWEDEN_API_BASE " , " http://fake-azure-sweden-api-base " )
2024-05-12 07:55:57 +08:00
monkeypatch . setenv ( " REDIS_HOST " , " localhost " )
Set fake env vars for `client_no_auth` fixture
This allows all of the tests in `test_proxy_server.py` to pass, with the
exception of `test_load_router_config`, without needing to set up real
environment variables.
Before:
```shell
$ env -i PATH=$PATH poetry run pytest litellm/tests/test_proxy_server.py -k 'not test_load_router_config' --disable-warnings
...
========================================================== short test summary info ===========================================================
ERROR litellm/tests/test_proxy_server.py::test_bedrock_embedding - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_chat_completion - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_chat_completion_azure - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_chat_completion_optional_params - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_embedding - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_engines_model_chat_completions - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_health - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_img_gen - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
ERROR litellm/tests/test_proxy_server.py::test_openai_deployments_model_chat_completions_azure - openai.OpenAIError: The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY enviro...
========================================== 2 skipped, 1 deselected, 39 warnings, 9 errors in 3.24s ===========================================
```
After:
```shell
$ env -i PATH=$PATH poetry run pytest litellm/tests/test_proxy_server.py -k 'not test_load_router_config' --disable-warnings
============================================================ test session starts =============================================================
platform darwin -- Python 3.12.3, pytest-7.4.4, pluggy-1.5.0
rootdir: /Users/abramowi/Code/OpenSource/litellm
plugins: anyio-4.3.0, asyncio-0.23.6, mock-3.14.0
asyncio: mode=Mode.STRICT
collected 12 items / 1 deselected / 11 selected
litellm/tests/test_proxy_server.py s.........s [100%]
========================================== 9 passed, 2 skipped, 1 deselected, 48 warnings in 8.42s ===========================================
```
2024-05-12 06:22:30 +08:00
@pytest.fixture ( scope = " function " )
def client_no_auth ( fake_env_vars ) :
2023-12-12 14:11:11 +08:00
# Assuming litellm.proxy.proxy_server is an object
from litellm . proxy . proxy_server import cleanup_router_config_variables
2023-12-25 16:40:38 +08:00
2023-12-12 14:11:11 +08:00
cleanup_router_config_variables ( )
2023-12-12 09:57:58 +08:00
filepath = os . path . dirname ( os . path . abspath ( __file__ ) )
2023-12-12 13:30:02 +08:00
config_fp = f " { filepath } /test_configs/test_config_no_auth.yaml "
2023-12-12 12:03:01 +08:00
# initialize can get run in parallel, it sets specific variables for the fast api app, sinc eit gets run in parallel different tests use the wrong variables
2024-01-04 20:58:18 +08:00
asyncio . run ( initialize ( config = config_fp , debug = True ) )
2023-12-12 09:57:58 +08:00
return TestClient ( app )
2023-12-25 16:40:38 +08:00
2024-05-03 04:36:23 +08:00
@mock_patch_acompletion ( )
def test_chat_completion ( mock_acompletion , client_no_auth ) :
2023-12-07 10:38:44 +08:00
global headers
2023-11-24 12:27:44 +08:00
try :
# Your test data
test_data = {
" model " : " gpt-3.5-turbo " ,
" messages " : [
2023-12-25 16:40:38 +08:00
{ " role " : " user " , " content " : " hi " } ,
2023-11-24 12:27:44 +08:00
] ,
" max_tokens " : 10 ,
}
2023-12-25 16:40:38 +08:00
2023-12-12 14:11:11 +08:00
print ( " testing proxy server with chat completions " )
2023-12-12 13:30:02 +08:00
response = client_no_auth . post ( " /v1/chat/completions " , json = test_data )
2024-05-03 04:36:23 +08:00
mock_acompletion . assert_called_once_with (
model = " gpt-3.5-turbo " ,
messages = [
{ " role " : " user " , " content " : " hi " } ,
] ,
max_tokens = 10 ,
litellm_call_id = mock . ANY ,
litellm_logging_obj = mock . ANY ,
request_timeout = mock . ANY ,
specific_deployment = True ,
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
2023-12-06 03:13:09 +08:00
print ( f " response - { response . text } " )
2023-11-24 12:27:44 +08:00
assert response . status_code == 200
result = response . json ( )
print ( f " Received response: { result } " )
except Exception as e :
2023-12-06 03:13:09 +08:00
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-11-24 12:27:44 +08:00
2023-12-25 16:40:38 +08:00
2024-07-21 09:39:05 +08:00
@mock_patch_acompletion ( )
@pytest.mark.asyncio
async def test_team_disable_guardrails ( mock_acompletion , client_no_auth ) :
"""
If team not allowed to turn on / off guardrails
Raise 403 forbidden error , if request is made by team on ` / key / generate ` or ` / chat / completions ` .
"""
import asyncio
import json
import time
from fastapi import HTTPException , Request
from starlette . datastructures import URL
2024-08-01 02:49:07 +08:00
from litellm . proxy . _types import (
LiteLLM_TeamTable ,
LiteLLM_TeamTableCachedObj ,
ProxyException ,
UserAPIKeyAuth ,
)
2024-07-21 09:39:05 +08:00
from litellm . proxy . auth . user_api_key_auth import user_api_key_auth
from litellm . proxy . proxy_server import hash_token , user_api_key_cache
_team_id = " 1234 "
user_key = " sk-12345678 "
valid_token = UserAPIKeyAuth (
team_id = _team_id ,
team_blocked = True ,
token = hash_token ( user_key ) ,
last_refreshed_at = time . time ( ) ,
)
await asyncio . sleep ( 1 )
2024-08-01 02:49:07 +08:00
team_obj = LiteLLM_TeamTableCachedObj (
2024-07-21 09:39:05 +08:00
team_id = _team_id ,
blocked = False ,
last_refreshed_at = time . time ( ) ,
metadata = { " guardrails " : { " modify_guardrails " : False } } ,
)
user_api_key_cache . set_cache ( key = hash_token ( user_key ) , value = valid_token )
user_api_key_cache . set_cache ( key = " team_id: {} " . format ( _team_id ) , value = team_obj )
setattr ( litellm . proxy . proxy_server , " user_api_key_cache " , user_api_key_cache )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
setattr ( litellm . proxy . proxy_server , " prisma_client " , " hello-world " )
request = Request ( scope = { " type " : " http " } )
request . _url = URL ( url = " /chat/completions " )
body = { " metadata " : { " guardrails " : { " hide_secrets " : False } } }
json_bytes = json . dumps ( body ) . encode ( " utf-8 " )
request . _body = json_bytes
try :
await user_api_key_auth ( request = request , api_key = " Bearer " + user_key )
pytest . fail ( " Expected to raise 403 forbidden error. " )
except ProxyException as e :
2024-08-01 02:49:07 +08:00
assert e . code == str ( 403 )
2024-07-21 09:39:05 +08:00
2024-09-29 04:28:01 +08:00
from tests . local_testing . test_custom_callback_input import CompletionCustomHandler
2024-07-17 08:15:20 +08:00
@mock_patch_acompletion ( )
def test_custom_logger_failure_handler ( mock_acompletion , client_no_auth ) :
from litellm . proxy . _types import UserAPIKeyAuth
from litellm . proxy . proxy_server import hash_token , user_api_key_cache
rpm_limit = 0
mock_api_key = " sk-my-test-key "
cache_value = UserAPIKeyAuth ( token = hash_token ( mock_api_key ) , rpm_limit = rpm_limit )
user_api_key_cache . set_cache ( key = hash_token ( mock_api_key ) , value = cache_value )
mock_logger = CustomLogger ( )
mock_logger_unit_tests = CompletionCustomHandler ( )
proxy_logging_obj : ProxyLogging = getattr (
litellm . proxy . proxy_server , " proxy_logging_obj "
)
litellm . callbacks = [ mock_logger , mock_logger_unit_tests ]
proxy_logging_obj . _init_litellm_callbacks ( llm_router = None )
setattr ( litellm . proxy . proxy_server , " user_api_key_cache " , user_api_key_cache )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
setattr ( litellm . proxy . proxy_server , " prisma_client " , " FAKE-VAR " )
setattr ( litellm . proxy . proxy_server , " proxy_logging_obj " , proxy_logging_obj )
with patch . object (
mock_logger , " async_log_failure_event " , new = AsyncMock ( )
) as mock_failed_alert :
# Your test data
test_data = {
" model " : " gpt-3.5-turbo " ,
" messages " : [
{ " role " : " user " , " content " : " hi " } ,
] ,
" max_tokens " : 10 ,
}
print ( " testing proxy server with chat completions " )
response = client_no_auth . post (
" /v1/chat/completions " ,
json = test_data ,
headers = { " Authorization " : " Bearer {} " . format ( mock_api_key ) } ,
)
assert response . status_code == 429
# confirm async_log_failure_event is called
mock_failed_alert . assert_called ( )
assert len ( mock_logger_unit_tests . errors ) == 0
2024-05-04 08:56:39 +08:00
@mock_patch_acompletion ( )
def test_engines_model_chat_completions ( mock_acompletion , client_no_auth ) :
global headers
try :
# Your test data
test_data = {
" model " : " gpt-3.5-turbo " ,
" messages " : [
{ " role " : " user " , " content " : " hi " } ,
] ,
" max_tokens " : 10 ,
}
print ( " testing proxy server with chat completions " )
2024-06-02 18:49:34 +08:00
response = client_no_auth . post (
" /engines/gpt-3.5-turbo/chat/completions " , json = test_data
)
2024-05-04 08:56:39 +08:00
mock_acompletion . assert_called_once_with (
model = " gpt-3.5-turbo " ,
messages = [
{ " role " : " user " , " content " : " hi " } ,
] ,
max_tokens = 10 ,
litellm_call_id = mock . ANY ,
litellm_logging_obj = mock . ANY ,
request_timeout = mock . ANY ,
specific_deployment = True ,
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
print ( f " response - { response . text } " )
assert response . status_code == 200
result = response . json ( )
print ( f " Received response: { result } " )
except Exception as e :
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-11-24 12:56:40 +08:00
2023-12-12 12:03:01 +08:00
2024-05-03 04:36:23 +08:00
@mock_patch_acompletion ( )
def test_chat_completion_azure ( mock_acompletion , client_no_auth ) :
2023-12-07 10:38:44 +08:00
global headers
2023-11-24 12:56:40 +08:00
try :
# Your test data
test_data = {
" model " : " azure/chatgpt-v-2 " ,
" messages " : [
2023-12-25 16:40:38 +08:00
{ " role " : " user " , " content " : " write 1 sentence poem " } ,
2023-11-24 12:56:40 +08:00
] ,
" max_tokens " : 10 ,
}
2023-12-25 16:40:38 +08:00
2023-12-12 14:11:11 +08:00
print ( " testing proxy server with Azure Request /chat/completions " )
2023-12-12 13:30:02 +08:00
response = client_no_auth . post ( " /v1/chat/completions " , json = test_data )
2023-11-24 12:56:40 +08:00
2024-05-03 04:36:23 +08:00
mock_acompletion . assert_called_once_with (
model = " azure/chatgpt-v-2 " ,
messages = [
{ " role " : " user " , " content " : " write 1 sentence poem " } ,
] ,
max_tokens = 10 ,
litellm_call_id = mock . ANY ,
litellm_logging_obj = mock . ANY ,
request_timeout = mock . ANY ,
specific_deployment = True ,
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
2023-11-24 12:56:40 +08:00
assert response . status_code == 200
result = response . json ( )
print ( f " Received response: { result } " )
2023-12-25 16:40:38 +08:00
assert len ( result [ " choices " ] [ 0 ] [ " message " ] [ " content " ] ) > 0
2023-11-24 12:56:40 +08:00
except Exception as e :
2023-12-06 03:13:09 +08:00
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-11-24 12:56:40 +08:00
2023-12-25 16:40:38 +08:00
2023-11-24 12:56:40 +08:00
# Run the test
2023-11-24 13:16:50 +08:00
# test_chat_completion_azure()
2023-11-24 12:56:40 +08:00
2023-12-25 16:40:38 +08:00
2024-05-03 03:24:49 +08:00
@mock_patch_acompletion ( )
2024-06-02 18:49:34 +08:00
def test_openai_deployments_model_chat_completions_azure (
mock_acompletion , client_no_auth
) :
2024-05-03 01:27:32 +08:00
global headers
try :
# Your test data
test_data = {
" model " : " azure/chatgpt-v-2 " ,
" messages " : [
{ " role " : " user " , " content " : " write 1 sentence poem " } ,
] ,
" max_tokens " : 10 ,
}
url = " /openai/deployments/azure/chatgpt-v-2/chat/completions "
print ( f " testing proxy server with Azure Request { url } " )
response = client_no_auth . post ( url , json = test_data )
2024-05-03 04:36:23 +08:00
mock_acompletion . assert_called_once_with (
model = " azure/chatgpt-v-2 " ,
messages = [
{ " role " : " user " , " content " : " write 1 sentence poem " } ,
] ,
max_tokens = 10 ,
litellm_call_id = mock . ANY ,
litellm_logging_obj = mock . ANY ,
request_timeout = mock . ANY ,
specific_deployment = True ,
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
2024-05-03 01:27:32 +08:00
assert response . status_code == 200
result = response . json ( )
print ( f " Received response: { result } " )
assert len ( result [ " choices " ] [ 0 ] [ " message " ] [ " content " ] ) > 0
except Exception as e :
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
# Run the test
# test_openai_deployments_model_chat_completions_azure()
2023-12-25 16:40:38 +08:00
### EMBEDDING
2024-05-03 04:36:23 +08:00
@mock_patch_aembedding ( )
def test_embedding ( mock_aembedding , client_no_auth ) :
2023-12-07 10:38:44 +08:00
global headers
2023-12-25 16:40:38 +08:00
from litellm . proxy . proxy_server import user_custom_auth
2023-12-12 14:11:11 +08:00
2023-11-24 13:16:50 +08:00
try :
test_data = {
" model " : " azure/azure-embedding-model " ,
" input " : [ " good morning from litellm " ] ,
}
2023-12-15 06:17:33 +08:00
response = client_no_auth . post ( " /v1/embeddings " , json = test_data )
2024-05-03 04:36:23 +08:00
mock_aembedding . assert_called_once_with (
model = " azure/azure-embedding-model " ,
input = [ " good morning from litellm " ] ,
specific_deployment = True ,
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
2023-12-15 06:17:33 +08:00
assert response . status_code == 200
result = response . json ( )
print ( len ( result [ " data " ] [ 0 ] [ " embedding " ] ) )
2023-12-25 16:40:38 +08:00
assert len ( result [ " data " ] [ 0 ] [ " embedding " ] ) > 10 # this usually has len==1536 so
2023-12-15 06:17:33 +08:00
except Exception as e :
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-12-25 16:40:38 +08:00
2024-05-03 04:36:23 +08:00
@mock_patch_aembedding ( )
def test_bedrock_embedding ( mock_aembedding , client_no_auth ) :
2023-12-15 06:17:33 +08:00
global headers
2023-12-25 16:40:38 +08:00
from litellm . proxy . proxy_server import user_custom_auth
2023-12-15 06:17:33 +08:00
try :
test_data = {
" model " : " amazon-embeddings " ,
" input " : [ " good morning from litellm " ] ,
}
2023-12-12 13:30:02 +08:00
response = client_no_auth . post ( " /v1/embeddings " , json = test_data )
2023-11-24 12:56:40 +08:00
2024-05-03 04:36:23 +08:00
mock_aembedding . assert_called_once_with (
model = " amazon-embeddings " ,
input = [ " good morning from litellm " ] ,
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
2023-11-24 13:16:50 +08:00
assert response . status_code == 200
result = response . json ( )
print ( len ( result [ " data " ] [ 0 ] [ " embedding " ] ) )
2023-12-25 16:40:38 +08:00
assert len ( result [ " data " ] [ 0 ] [ " embedding " ] ) > 10 # this usually has len==1536 so
2023-11-24 13:16:50 +08:00
except Exception as e :
2023-12-06 03:13:09 +08:00
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-11-24 12:56:40 +08:00
2023-12-25 16:40:38 +08:00
2024-02-29 11:16:11 +08:00
@pytest.mark.skip ( reason = " AWS Suspended Account " )
2023-12-15 06:38:27 +08:00
def test_sagemaker_embedding ( client_no_auth ) :
global headers
2023-12-25 16:40:38 +08:00
from litellm . proxy . proxy_server import user_custom_auth
2023-12-15 06:38:27 +08:00
try :
test_data = {
" model " : " GPT-J 6B - Sagemaker Text Embedding (Internal) " ,
" input " : [ " good morning from litellm " ] ,
}
response = client_no_auth . post ( " /v1/embeddings " , json = test_data )
assert response . status_code == 200
result = response . json ( )
print ( len ( result [ " data " ] [ 0 ] [ " embedding " ] ) )
2023-12-25 16:40:38 +08:00
assert len ( result [ " data " ] [ 0 ] [ " embedding " ] ) > 10 # this usually has len==1536 so
2023-12-15 06:38:27 +08:00
except Exception as e :
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-12-25 16:40:38 +08:00
2023-11-24 13:16:50 +08:00
# Run the test
2023-11-24 12:56:40 +08:00
# test_embedding()
2023-12-21 18:09:09 +08:00
#### IMAGE GENERATION
2023-12-25 16:40:38 +08:00
2024-05-03 04:36:23 +08:00
@mock_patch_aimage_generation ( )
def test_img_gen ( mock_aimage_generation , client_no_auth ) :
2023-12-21 18:09:09 +08:00
global headers
2023-12-25 16:40:38 +08:00
from litellm . proxy . proxy_server import user_custom_auth
2023-12-21 18:09:09 +08:00
try :
test_data = {
" model " : " dall-e-3 " ,
" prompt " : " A cute baby sea otter " ,
" n " : 1 ,
2023-12-25 16:40:38 +08:00
" size " : " 1024x1024 " ,
2023-12-21 18:09:09 +08:00
}
response = client_no_auth . post ( " /v1/images/generations " , json = test_data )
2024-05-03 04:36:23 +08:00
mock_aimage_generation . assert_called_once_with (
2024-06-02 18:49:34 +08:00
model = " dall-e-3 " ,
prompt = " A cute baby sea otter " ,
2024-05-03 04:36:23 +08:00
n = 1 ,
2024-06-02 18:49:34 +08:00
size = " 1024x1024 " ,
2024-05-03 04:36:23 +08:00
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
2023-12-21 18:09:09 +08:00
assert response . status_code == 200
result = response . json ( )
print ( len ( result [ " data " ] [ 0 ] [ " url " ] ) )
2023-12-25 16:40:38 +08:00
assert len ( result [ " data " ] [ 0 ] [ " url " ] ) > 10
2023-12-21 18:09:09 +08:00
except Exception as e :
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-12-03 06:15:38 +08:00
2023-12-25 16:40:38 +08:00
#### ADDITIONAL
2024-04-04 13:37:51 +08:00
@pytest.mark.skip ( reason = " test via docker tests. Requires prisma client. " )
2023-12-12 13:30:02 +08:00
def test_add_new_model ( client_no_auth ) :
2023-12-07 10:38:44 +08:00
global headers
2023-12-25 16:40:38 +08:00
try :
2023-12-03 06:15:38 +08:00
test_data = {
" model_name " : " test_openai_models " ,
" litellm_params " : {
2023-12-25 16:40:38 +08:00
" model " : " gpt-3.5-turbo " ,
2023-12-03 06:15:38 +08:00
} ,
2023-12-25 16:40:38 +08:00
" model_info " : { " description " : " this is a test openai model " } ,
2023-12-03 06:15:38 +08:00
}
2023-12-12 13:30:02 +08:00
client_no_auth . post ( " /model/new " , json = test_data , headers = headers )
response = client_no_auth . get ( " /model/info " , headers = headers )
2023-12-03 06:15:38 +08:00
assert response . status_code == 200
2023-12-25 16:40:38 +08:00
result = response . json ( )
2023-12-03 06:15:38 +08:00
print ( f " response: { result } " )
model_info = None
for m in result [ " data " ] :
2023-12-10 14:43:27 +08:00
if m [ " model_name " ] == " test_openai_models " :
model_info = m [ " model_info " ]
2023-12-03 06:15:38 +08:00
assert model_info [ " description " ] == " this is a test openai model "
2023-12-25 16:40:38 +08:00
except Exception as e :
2023-12-03 06:15:38 +08:00
pytest . fail ( f " LiteLLM Proxy test failed. Exception { str ( e ) } " )
2023-12-22 13:38:44 +08:00
def test_health ( client_no_auth ) :
global headers
2024-07-17 08:15:20 +08:00
import logging
2023-12-22 13:38:44 +08:00
import time
2024-07-17 08:15:20 +08:00
2024-06-16 06:09:49 +08:00
from litellm . _logging import verbose_logger , verbose_proxy_logger
verbose_proxy_logger . setLevel ( logging . DEBUG )
2023-12-25 16:40:38 +08:00
2023-12-22 13:38:44 +08:00
try :
response = client_no_auth . get ( " /health " )
assert response . status_code == 200
except Exception as e :
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
2023-12-25 16:40:38 +08:00
2023-12-05 02:19:35 +08:00
# test_add_new_model()
from litellm . integrations . custom_logger import CustomLogger
2023-12-25 16:40:38 +08:00
2023-12-05 02:19:35 +08:00
class MyCustomHandler ( CustomLogger ) :
2023-12-25 16:40:38 +08:00
def log_pre_api_call ( self , model , messages , kwargs ) :
2023-12-05 02:19:35 +08:00
print ( f " Pre-API Call " )
2023-12-25 16:40:38 +08:00
def log_success_event ( self , kwargs , response_obj , start_time , end_time ) :
2023-12-05 02:19:35 +08:00
print ( f " On Success " )
assert kwargs [ " user " ] == " proxy-user "
assert kwargs [ " model " ] == " gpt-3.5-turbo "
assert kwargs [ " max_tokens " ] == 10
2023-12-25 16:40:38 +08:00
2023-12-05 02:19:35 +08:00
customHandler = MyCustomHandler ( )
2024-05-03 04:36:23 +08:00
@mock_patch_acompletion ( )
def test_chat_completion_optional_params ( mock_acompletion , client_no_auth ) :
2023-12-05 02:19:35 +08:00
# [PROXY: PROD TEST] - DO NOT DELETE
# This tests if all the /chat/completion params are passed to litellm
try :
# Your test data
2023-12-25 16:40:38 +08:00
litellm . set_verbose = True
2023-12-05 02:19:35 +08:00
test_data = {
" model " : " gpt-3.5-turbo " ,
" messages " : [
2023-12-25 16:40:38 +08:00
{ " role " : " user " , " content " : " hi " } ,
2023-12-05 02:19:35 +08:00
] ,
" max_tokens " : 10 ,
2023-12-25 16:40:38 +08:00
" user " : " proxy-user " ,
2023-12-05 02:19:35 +08:00
}
2023-12-25 16:40:38 +08:00
2023-12-05 02:19:35 +08:00
litellm . callbacks = [ customHandler ]
print ( " testing proxy server: optional params " )
2023-12-12 13:30:02 +08:00
response = client_no_auth . post ( " /v1/chat/completions " , json = test_data )
2024-05-03 04:36:23 +08:00
mock_acompletion . assert_called_once_with (
model = " gpt-3.5-turbo " ,
messages = [
{ " role " : " user " , " content " : " hi " } ,
] ,
max_tokens = 10 ,
user = " proxy-user " ,
litellm_call_id = mock . ANY ,
litellm_logging_obj = mock . ANY ,
request_timeout = mock . ANY ,
specific_deployment = True ,
metadata = mock . ANY ,
proxy_server_request = mock . ANY ,
)
2023-12-05 02:19:35 +08:00
assert response . status_code == 200
result = response . json ( )
print ( f " Received response: { result } " )
except Exception as e :
pytest . fail ( " LiteLLM Proxy test failed. Exception " , e )
2023-12-25 16:40:38 +08:00
2023-12-05 02:19:35 +08:00
# Run the test
2023-12-05 05:16:19 +08:00
# test_chat_completion_optional_params()
2024-07-27 10:03:42 +08:00
2023-12-25 16:40:38 +08:00
# Test Reading config.yaml file
2024-01-04 18:41:23 +08:00
from litellm . proxy . proxy_server import ProxyConfig
2023-12-05 05:16:19 +08:00
2023-12-25 16:40:38 +08:00
2024-08-25 10:32:22 +08:00
@pytest.mark.skip ( reason = " local variable conflicts. needs to be refactored. " )
2024-05-12 07:55:57 +08:00
@mock.patch ( " litellm.proxy.proxy_server.litellm.Cache " )
def test_load_router_config ( mock_cache , fake_env_vars ) :
mock_cache . return_value . cache . __dict__ = { " redis_client " : None }
mock_cache . return_value . supported_call_types = [
" completion " ,
" acompletion " ,
" embedding " ,
" aembedding " ,
" atranscription " ,
" transcription " ,
]
2023-12-05 05:16:19 +08:00
try :
2024-01-05 16:07:31 +08:00
import asyncio
2023-12-05 05:16:19 +08:00
print ( " testing reading config " )
2023-12-05 05:24:35 +08:00
# this is a basic config.yaml with only a model
2023-12-05 07:24:46 +08:00
filepath = os . path . dirname ( os . path . abspath ( __file__ ) )
2024-01-04 18:41:23 +08:00
proxy_config = ProxyConfig ( )
2024-01-05 16:07:31 +08:00
result = asyncio . run (
proxy_config . load_config (
router = None ,
config_file_path = f " { filepath } /example_config_yaml/simple_config.yaml " ,
)
2023-12-25 16:40:38 +08:00
)
2023-12-05 05:16:19 +08:00
print ( result )
assert len ( result [ 1 ] ) == 1
2023-12-05 05:24:35 +08:00
# this is a load balancing config yaml
2024-01-05 16:07:31 +08:00
result = asyncio . run (
proxy_config . load_config (
router = None ,
config_file_path = f " { filepath } /example_config_yaml/azure_config.yaml " ,
)
2023-12-25 16:40:38 +08:00
)
2023-12-05 05:24:35 +08:00
print ( result )
assert len ( result [ 1 ] ) == 2
2023-12-05 06:49:59 +08:00
# config with general settings - custom callbacks
2024-01-05 16:07:31 +08:00
result = asyncio . run (
proxy_config . load_config (
router = None ,
config_file_path = f " { filepath } /example_config_yaml/azure_config.yaml " ,
)
2023-12-25 16:40:38 +08:00
)
2023-12-05 06:49:59 +08:00
print ( result )
assert len ( result [ 1 ] ) == 2
2023-12-05 05:24:35 +08:00
2023-12-16 17:15:06 +08:00
# tests for litellm.cache set from config
print ( " testing reading proxy config for cache " )
litellm . cache = None
2024-01-05 16:07:31 +08:00
asyncio . run (
proxy_config . load_config (
router = None ,
config_file_path = f " { filepath } /example_config_yaml/cache_no_params.yaml " ,
)
2023-12-16 17:15:06 +08:00
)
assert litellm . cache is not None
2023-12-25 16:40:38 +08:00
assert " redis_client " in vars (
litellm . cache . cache
) # it should default to redis on proxy
assert litellm . cache . supported_call_types == [
" completion " ,
" acompletion " ,
" embedding " ,
" aembedding " ,
2024-03-10 10:47:20 +08:00
" atranscription " ,
" transcription " ,
2023-12-25 16:40:38 +08:00
] # init with all call types
2024-01-05 16:07:31 +08:00
litellm . disable_cache ( )
2023-12-16 17:15:06 +08:00
print ( " testing reading proxy config for cache with params " )
2024-05-12 07:55:57 +08:00
mock_cache . return_value . supported_call_types = [
" embedding " ,
" aembedding " ,
]
2024-01-05 16:07:31 +08:00
asyncio . run (
proxy_config . load_config (
router = None ,
config_file_path = f " { filepath } /example_config_yaml/cache_with_params.yaml " ,
)
2023-12-16 17:15:06 +08:00
)
assert litellm . cache is not None
print ( litellm . cache )
print ( litellm . cache . supported_call_types )
print ( vars ( litellm . cache . cache ) )
2023-12-25 16:40:38 +08:00
assert " redis_client " in vars (
litellm . cache . cache
) # it should default to redis on proxy
assert litellm . cache . supported_call_types == [
" embedding " ,
" aembedding " ,
] # init with all call types
2023-12-16 17:15:06 +08:00
2023-12-05 05:16:19 +08:00
except Exception as e :
2024-04-20 10:22:24 +08:00
pytest . fail (
f " Proxy: Got exception reading config: { str ( e ) } \n { traceback . format_exc ( ) } "
)
2023-12-25 16:40:38 +08:00
# test_load_router_config()
2024-07-25 09:14:49 +08:00
@pytest.mark.asyncio
async def test_team_update_redis ( ) :
"""
Tests if team update , updates the redis cache if set
"""
from litellm . caching import DualCache , RedisCache
2024-08-01 02:49:07 +08:00
from litellm . proxy . _types import LiteLLM_TeamTableCachedObj
2024-07-25 09:14:49 +08:00
from litellm . proxy . auth . auth_checks import _cache_team_object
proxy_logging_obj : ProxyLogging = getattr (
litellm . proxy . proxy_server , " proxy_logging_obj "
)
2024-09-26 10:56:17 +08:00
redis_cache = RedisCache ( )
2024-07-25 09:14:49 +08:00
with patch . object (
2024-09-26 10:56:17 +08:00
redis_cache ,
2024-07-25 09:14:49 +08:00
" async_set_cache " ,
2024-08-08 06:37:02 +08:00
new = AsyncMock ( ) ,
2024-07-25 09:14:49 +08:00
) as mock_client :
await _cache_team_object (
team_id = " 1234 " ,
2024-09-15 01:02:55 +08:00
team_table = LiteLLM_TeamTableCachedObj ( team_id = " 1234 " ) ,
2024-09-26 10:56:17 +08:00
user_api_key_cache = DualCache ( redis_cache = redis_cache ) ,
2024-07-25 09:14:49 +08:00
proxy_logging_obj = proxy_logging_obj ,
)
2024-08-08 06:37:02 +08:00
mock_client . assert_called ( )
2024-07-25 09:14:49 +08:00
@pytest.mark.asyncio
async def test_get_team_redis ( client_no_auth ) :
"""
Tests if get_team_object gets value from redis cache , if set
"""
from litellm . caching import DualCache , RedisCache
2024-09-15 01:02:55 +08:00
from litellm . proxy . auth . auth_checks import get_team_object
2024-07-25 09:14:49 +08:00
proxy_logging_obj : ProxyLogging = getattr (
litellm . proxy . proxy_server , " proxy_logging_obj "
)
2024-09-26 10:56:17 +08:00
redis_cache = RedisCache ( )
2024-07-25 09:14:49 +08:00
with patch . object (
2024-09-26 10:56:17 +08:00
redis_cache ,
2024-07-25 09:14:49 +08:00
" async_get_cache " ,
new = AsyncMock ( ) ,
) as mock_client :
try :
await get_team_object (
team_id = " 1234 " ,
2024-09-26 10:56:17 +08:00
user_api_key_cache = DualCache ( redis_cache = redis_cache ) ,
2024-07-25 09:14:49 +08:00
parent_otel_span = None ,
proxy_logging_obj = proxy_logging_obj ,
2024-08-08 06:37:02 +08:00
prisma_client = AsyncMock ( ) ,
2024-07-25 09:14:49 +08:00
)
except Exception as e :
pass
mock_client . assert_called_once ( )
2024-08-08 09:50:40 +08:00
import random
import uuid
2024-08-09 10:14:43 +08:00
from unittest . mock import AsyncMock , MagicMock , PropertyMock , patch
2024-08-08 09:50:40 +08:00
2024-08-09 10:14:43 +08:00
from litellm . proxy . _types import (
LitellmUserRoles ,
NewUserRequest ,
TeamMemberAddRequest ,
UserAPIKeyAuth ,
)
2024-08-08 09:50:40 +08:00
from litellm . proxy . management_endpoints . internal_user_endpoints import new_user
2024-08-09 10:14:43 +08:00
from litellm . proxy . management_endpoints . team_endpoints import team_member_add
2024-09-29 04:28:01 +08:00
from tests . local_testing . test_key_generate_prisma import prisma_client
2024-08-08 09:50:40 +08:00
2024-08-09 08:59:30 +08:00
@pytest.mark.parametrize (
" user_role " ,
[ LitellmUserRoles . INTERNAL_USER . value , LitellmUserRoles . PROXY_ADMIN . value ] ,
)
2024-08-08 09:50:40 +08:00
@pytest.mark.asyncio
2024-08-09 08:59:30 +08:00
async def test_create_user_default_budget ( prisma_client , user_role ) :
2024-08-08 09:50:40 +08:00
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
2024-08-08 09:59:46 +08:00
setattr ( litellm , " max_internal_user_budget " , 10 )
2024-08-09 08:59:30 +08:00
setattr ( litellm , " internal_user_budget_duration " , " 5m " )
2024-08-08 09:50:40 +08:00
await litellm . proxy . proxy_server . prisma_client . connect ( )
user = f " ishaan { uuid . uuid4 ( ) . hex } "
2024-08-09 08:59:30 +08:00
request = NewUserRequest (
user_id = user , user_role = user_role
) # create a key with no budget
2024-08-08 09:50:40 +08:00
with patch . object (
litellm . proxy . proxy_server . prisma_client , " insert_data " , new = AsyncMock ( )
) as mock_client :
await new_user (
request ,
)
mock_client . assert_called ( )
print ( f " mock_client.call_args: { mock_client . call_args } " )
print ( " mock_client.call_args.kwargs: {} " . format ( mock_client . call_args . kwargs ) )
2024-08-09 08:59:30 +08:00
if user_role == LitellmUserRoles . INTERNAL_USER . value :
assert (
mock_client . call_args . kwargs [ " data " ] [ " max_budget " ]
== litellm . max_internal_user_budget
)
assert (
mock_client . call_args . kwargs [ " data " ] [ " budget_duration " ]
== litellm . internal_user_budget_duration
)
else :
assert mock_client . call_args . kwargs [ " data " ] [ " max_budget " ] is None
assert mock_client . call_args . kwargs [ " data " ] [ " budget_duration " ] is None
2024-08-09 10:14:43 +08:00
@pytest.mark.parametrize ( " new_member_method " , [ " user_id " , " user_email " ] )
@pytest.mark.asyncio
async def test_create_team_member_add ( prisma_client , new_member_method ) :
import time
2024-08-10 03:15:45 +08:00
from fastapi import Request
2024-08-25 10:32:22 +08:00
from litellm . proxy . _types import LiteLLM_TeamTableCachedObj , LiteLLM_UserTable
2024-08-09 10:14:43 +08:00
from litellm . proxy . proxy_server import hash_token , user_api_key_cache
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
setattr ( litellm , " max_internal_user_budget " , 10 )
setattr ( litellm , " internal_user_budget_duration " , " 5m " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
user = f " ishaan { uuid . uuid4 ( ) . hex } "
_team_id = " litellm-test-client-id-new "
team_obj = LiteLLM_TeamTableCachedObj (
team_id = _team_id ,
blocked = False ,
last_refreshed_at = time . time ( ) ,
metadata = { " guardrails " : { " modify_guardrails " : False } } ,
)
# user_api_key_cache.set_cache(key=hash_token(user_key), value=valid_token)
user_api_key_cache . set_cache ( key = " team_id: {} " . format ( _team_id ) , value = team_obj )
setattr ( litellm . proxy . proxy_server , " user_api_key_cache " , user_api_key_cache )
if new_member_method == " user_id " :
data = {
" team_id " : _team_id ,
" member " : [ { " role " : " user " , " user_id " : user } ] ,
}
elif new_member_method == " user_email " :
data = {
" team_id " : _team_id ,
" member " : [ { " role " : " user " , " user_email " : user } ] ,
}
team_member_add_request = TeamMemberAddRequest ( * * data )
with patch (
" litellm.proxy.proxy_server.prisma_client.db.litellm_usertable " ,
new_callable = AsyncMock ,
) as mock_litellm_usertable :
2024-08-25 10:32:22 +08:00
mock_client = AsyncMock (
return_value = LiteLLM_UserTable (
user_id = " 1234 " , max_budget = 100 , user_email = " 1234 "
)
)
2024-08-09 10:14:43 +08:00
mock_litellm_usertable . upsert = mock_client
mock_litellm_usertable . find_many = AsyncMock ( return_value = None )
2024-08-25 10:32:22 +08:00
team_mock_client = AsyncMock ( )
original_val = getattr (
litellm . proxy . proxy_server . prisma_client . db , " litellm_teamtable "
)
litellm . proxy . proxy_server . prisma_client . db . litellm_teamtable = team_mock_client
2024-09-15 01:02:55 +08:00
team_mock_client . update = AsyncMock (
return_value = LiteLLM_TeamTableCachedObj ( team_id = " 1234 " )
)
2024-08-09 10:14:43 +08:00
await team_member_add (
2024-08-10 03:15:45 +08:00
data = team_member_add_request ,
2024-08-21 23:37:04 +08:00
user_api_key_dict = UserAPIKeyAuth ( user_role = " proxy_admin " ) ,
2024-08-10 03:15:45 +08:00
http_request = Request (
scope = { " type " : " http " , " path " : " /user/new " } ,
) ,
2024-08-09 10:14:43 +08:00
)
mock_client . assert_called ( )
print ( f " mock_client.call_args: { mock_client . call_args } " )
print ( " mock_client.call_args.kwargs: {} " . format ( mock_client . call_args . kwargs ) )
assert (
mock_client . call_args . kwargs [ " data " ] [ " create " ] [ " max_budget " ]
== litellm . max_internal_user_budget
)
assert (
mock_client . call_args . kwargs [ " data " ] [ " create " ] [ " budget_duration " ]
== litellm . internal_user_budget_duration
)
2024-08-13 09:47:25 +08:00
2024-08-25 10:32:22 +08:00
litellm . proxy . proxy_server . prisma_client . db . litellm_teamtable = original_val
2024-08-13 09:47:25 +08:00
2024-08-21 05:01:12 +08:00
@pytest.mark.parametrize ( " team_member_role " , [ " admin " , " user " ] )
2024-08-21 07:25:13 +08:00
@pytest.mark.parametrize ( " team_route " , [ " /team/member_add " , " /team/member_delete " ] )
2024-08-21 05:01:12 +08:00
@pytest.mark.asyncio
async def test_create_team_member_add_team_admin_user_api_key_auth (
2024-08-21 07:25:13 +08:00
prisma_client , team_member_role , team_route
2024-08-21 05:01:12 +08:00
) :
import time
from fastapi import Request
from litellm . proxy . _types import LiteLLM_TeamTableCachedObj , Member
from litellm . proxy . proxy_server import (
ProxyException ,
hash_token ,
user_api_key_auth ,
user_api_key_cache ,
)
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
setattr ( litellm , " max_internal_user_budget " , 10 )
setattr ( litellm , " internal_user_budget_duration " , " 5m " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
user = f " ishaan { uuid . uuid4 ( ) . hex } "
_team_id = " litellm-test-client-id-new "
user_key = " sk-12345678 "
valid_token = UserAPIKeyAuth (
team_id = _team_id ,
token = hash_token ( user_key ) ,
team_member = Member ( role = team_member_role , user_id = user ) ,
last_refreshed_at = time . time ( ) ,
)
user_api_key_cache . set_cache ( key = hash_token ( user_key ) , value = valid_token )
team_obj = LiteLLM_TeamTableCachedObj (
team_id = _team_id ,
blocked = False ,
last_refreshed_at = time . time ( ) ,
metadata = { " guardrails " : { " modify_guardrails " : False } } ,
)
user_api_key_cache . set_cache ( key = " team_id: {} " . format ( _team_id ) , value = team_obj )
setattr ( litellm . proxy . proxy_server , " user_api_key_cache " , user_api_key_cache )
## TEST IF TEAM ADMIN ALLOWED TO CALL /MEMBER_ADD ENDPOINT
import json
from starlette . datastructures import URL
request = Request ( scope = { " type " : " http " } )
2024-08-21 07:25:13 +08:00
request . _url = URL ( url = team_route )
2024-08-21 05:01:12 +08:00
body = { }
json_bytes = json . dumps ( body ) . encode ( " utf-8 " )
request . _body = json_bytes
2024-08-21 07:57:18 +08:00
## ALLOWED BY USER_API_KEY_AUTH
await user_api_key_auth ( request = request , api_key = " Bearer " + user_key )
2024-08-21 05:01:12 +08:00
@pytest.mark.parametrize ( " new_member_method " , [ " user_id " , " user_email " ] )
2024-08-21 07:57:18 +08:00
@pytest.mark.parametrize ( " user_role " , [ " admin " , " user " ] )
2024-08-21 05:01:12 +08:00
@pytest.mark.asyncio
2024-08-21 07:57:18 +08:00
async def test_create_team_member_add_team_admin (
prisma_client , new_member_method , user_role
) :
2024-08-21 05:01:12 +08:00
"""
Relevant issue - https : / / github . com / BerriAI / litellm / issues / 5300
Allow team admins to :
- Add and remove team members
- raise error if team member not an existing ' internal_user '
"""
import time
from fastapi import Request
2024-08-25 10:32:22 +08:00
from litellm . proxy . _types import (
LiteLLM_TeamTableCachedObj ,
LiteLLM_UserTable ,
Member ,
)
2024-08-21 05:01:12 +08:00
from litellm . proxy . proxy_server import (
2024-08-21 07:57:18 +08:00
HTTPException ,
2024-08-21 05:01:12 +08:00
ProxyException ,
hash_token ,
user_api_key_auth ,
user_api_key_cache ,
)
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
setattr ( litellm , " max_internal_user_budget " , 10 )
setattr ( litellm , " internal_user_budget_duration " , " 5m " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
user = f " ishaan { uuid . uuid4 ( ) . hex } "
_team_id = " litellm-test-client-id-new "
user_key = " sk-12345678 "
valid_token = UserAPIKeyAuth (
team_id = _team_id ,
2024-08-21 07:57:18 +08:00
user_id = user ,
2024-08-21 05:01:12 +08:00
token = hash_token ( user_key ) ,
last_refreshed_at = time . time ( ) ,
)
user_api_key_cache . set_cache ( key = hash_token ( user_key ) , value = valid_token )
team_obj = LiteLLM_TeamTableCachedObj (
team_id = _team_id ,
blocked = False ,
last_refreshed_at = time . time ( ) ,
2024-08-21 07:57:18 +08:00
members_with_roles = [ Member ( role = user_role , user_id = user ) ] ,
2024-08-21 05:01:12 +08:00
metadata = { " guardrails " : { " modify_guardrails " : False } } ,
)
user_api_key_cache . set_cache ( key = " team_id: {} " . format ( _team_id ) , value = team_obj )
setattr ( litellm . proxy . proxy_server , " user_api_key_cache " , user_api_key_cache )
if new_member_method == " user_id " :
data = {
" team_id " : _team_id ,
" member " : [ { " role " : " user " , " user_id " : user } ] ,
}
elif new_member_method == " user_email " :
data = {
" team_id " : _team_id ,
" member " : [ { " role " : " user " , " user_email " : user } ] ,
}
team_member_add_request = TeamMemberAddRequest ( * * data )
with patch (
" litellm.proxy.proxy_server.prisma_client.db.litellm_usertable " ,
new_callable = AsyncMock ,
) as mock_litellm_usertable :
2024-08-25 10:32:22 +08:00
mock_client = AsyncMock (
return_value = LiteLLM_UserTable (
user_id = " 1234 " , max_budget = 100 , user_email = " 1234 "
)
)
2024-08-21 05:01:12 +08:00
mock_litellm_usertable . upsert = mock_client
mock_litellm_usertable . find_many = AsyncMock ( return_value = None )
2024-08-25 10:32:22 +08:00
team_mock_client = AsyncMock ( )
original_val = getattr (
litellm . proxy . proxy_server . prisma_client . db , " litellm_teamtable "
)
litellm . proxy . proxy_server . prisma_client . db . litellm_teamtable = team_mock_client
2024-09-15 01:02:55 +08:00
team_mock_client . update = AsyncMock (
return_value = LiteLLM_TeamTableCachedObj ( team_id = " 1234 " )
)
2024-08-25 10:32:22 +08:00
2024-08-21 07:57:18 +08:00
try :
await team_member_add (
data = team_member_add_request ,
user_api_key_dict = valid_token ,
http_request = Request (
scope = { " type " : " http " , " path " : " /user/new " } ,
) ,
)
except HTTPException as e :
if user_role == " user " :
assert e . status_code == 403
else :
raise e
2024-08-21 05:01:12 +08:00
mock_client . assert_called ( )
print ( f " mock_client.call_args: { mock_client . call_args } " )
print ( " mock_client.call_args.kwargs: {} " . format ( mock_client . call_args . kwargs ) )
assert (
mock_client . call_args . kwargs [ " data " ] [ " create " ] [ " max_budget " ]
== litellm . max_internal_user_budget
)
assert (
mock_client . call_args . kwargs [ " data " ] [ " create " ] [ " budget_duration " ]
== litellm . internal_user_budget_duration
)
2024-08-25 10:32:22 +08:00
litellm . proxy . proxy_server . prisma_client . db . litellm_teamtable = original_val
2024-08-21 05:01:12 +08:00
2024-08-13 09:47:25 +08:00
@pytest.mark.asyncio
async def test_user_info_team_list ( prisma_client ) :
""" Assert user_info for admin calls team_list function """
from litellm . proxy . _types import LiteLLM_UserTable
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
from litellm . proxy . management_endpoints . internal_user_endpoints import user_info
with patch (
" litellm.proxy.management_endpoints.team_endpoints.list_team " ,
new_callable = AsyncMock ,
) as mock_client :
prisma_client . get_data = AsyncMock (
return_value = LiteLLM_UserTable (
user_role = " proxy_admin " ,
user_id = " default_user_id " ,
max_budget = None ,
user_email = " " ,
)
)
try :
await user_info (
user_id = None ,
user_api_key_dict = UserAPIKeyAuth (
api_key = " sk-1234 " , user_id = " default_user_id "
) ,
)
except Exception :
pass
mock_client . assert_called ( )
2024-08-13 12:21:40 +08:00
2024-08-14 12:27:59 +08:00
2024-08-14 12:36:16 +08:00
@pytest.mark.skip ( reason = " Local test " )
2024-08-13 12:21:40 +08:00
@pytest.mark.asyncio
async def test_add_callback_via_key ( prisma_client ) :
"""
Test if callback specified in key , is used .
"""
global headers
import json
from fastapi import HTTPException , Request , Response
from starlette . datastructures import URL
from litellm . proxy . proxy_server import chat_completion
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
litellm . set_verbose = True
try :
# Your test data
test_data = {
" model " : " azure/chatgpt-v-2 " ,
" messages " : [
{ " role " : " user " , " content " : " write 1 sentence poem " } ,
] ,
" max_tokens " : 10 ,
" mock_response " : " Hello world " ,
" api_key " : " my-fake-key " ,
}
request = Request ( scope = { " type " : " http " , " method " : " POST " , " headers " : { } } )
request . _url = URL ( url = " /chat/completions " )
json_bytes = json . dumps ( test_data ) . encode ( " utf-8 " )
request . _body = json_bytes
with patch . object (
litellm . litellm_core_utils . litellm_logging ,
" LangFuseLogger " ,
new = MagicMock ( ) ,
) as mock_client :
resp = await chat_completion (
request = request ,
fastapi_response = Response ( ) ,
user_api_key_dict = UserAPIKeyAuth (
metadata = {
" logging " : [
{
" callback_name " : " langfuse " , # 'otel', 'langfuse', 'lunary'
" callback_type " : " success " , # set, if required by integration - future improvement, have logging tools work for success + failure by default
" callback_vars " : {
" langfuse_public_key " : " os.environ/LANGFUSE_PUBLIC_KEY " ,
" langfuse_secret_key " : " os.environ/LANGFUSE_SECRET_KEY " ,
" langfuse_host " : " https://us.cloud.langfuse.com " ,
} ,
}
]
}
) ,
)
print ( resp )
mock_client . assert_called ( )
mock_client . return_value . log_event . assert_called ( )
args , kwargs = mock_client . return_value . log_event . call_args
kwargs = kwargs [ " kwargs " ]
assert " user_api_key_metadata " in kwargs [ " litellm_params " ] [ " metadata " ]
assert (
" logging "
in kwargs [ " litellm_params " ] [ " metadata " ] [ " user_api_key_metadata " ]
)
checked_keys = False
for item in kwargs [ " litellm_params " ] [ " metadata " ] [ " user_api_key_metadata " ] [
" logging "
] :
for k , v in item [ " callback_vars " ] . items ( ) :
print ( " k= {} , v= {} " . format ( k , v ) )
if " key " in k :
assert " os.environ " in v
checked_keys = True
assert checked_keys
except Exception as e :
2024-08-14 12:27:59 +08:00
pytest . fail ( f " LiteLLM Proxy test failed. Exception - { str ( e ) } " )
@pytest.mark.asyncio
2024-09-10 07:44:37 +08:00
@pytest.mark.parametrize (
" callback_type, expected_success_callbacks, expected_failure_callbacks " ,
[
( " success " , [ " langfuse " ] , [ ] ) ,
( " failure " , [ ] , [ " langfuse " ] ) ,
( " success_and_failure " , [ " langfuse " ] , [ " langfuse " ] ) ,
] ,
)
async def test_add_callback_via_key_litellm_pre_call_utils (
prisma_client , callback_type , expected_success_callbacks , expected_failure_callbacks
) :
2024-08-14 12:27:59 +08:00
import json
from fastapi import HTTPException , Request , Response
from starlette . datastructures import URL
from litellm . proxy . litellm_pre_call_utils import add_litellm_data_to_request
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
proxy_config = getattr ( litellm . proxy . proxy_server , " proxy_config " )
request = Request ( scope = { " type " : " http " , " method " : " POST " , " headers " : { } } )
request . _url = URL ( url = " /chat/completions " )
test_data = {
" model " : " azure/chatgpt-v-2 " ,
" messages " : [
{ " role " : " user " , " content " : " write 1 sentence poem " } ,
] ,
" max_tokens " : 10 ,
" mock_response " : " Hello world " ,
" api_key " : " my-fake-key " ,
}
json_bytes = json . dumps ( test_data ) . encode ( " utf-8 " )
request . _body = json_bytes
data = {
" data " : {
" model " : " azure/chatgpt-v-2 " ,
" messages " : [ { " role " : " user " , " content " : " write 1 sentence poem " } ] ,
" max_tokens " : 10 ,
" mock_response " : " Hello world " ,
" api_key " : " my-fake-key " ,
} ,
" request " : request ,
" user_api_key_dict " : UserAPIKeyAuth (
token = None ,
key_name = None ,
key_alias = None ,
spend = 0.0 ,
max_budget = None ,
expires = None ,
models = [ ] ,
aliases = { } ,
config = { } ,
user_id = None ,
team_id = None ,
max_parallel_requests = None ,
metadata = {
" logging " : [
{
" callback_name " : " langfuse " ,
2024-09-10 07:44:37 +08:00
" callback_type " : callback_type ,
2024-08-14 12:27:59 +08:00
" callback_vars " : {
2024-08-22 04:36:33 +08:00
" langfuse_public_key " : " my-mock-public-key " ,
" langfuse_secret_key " : " my-mock-secret-key " ,
2024-08-14 12:27:59 +08:00
" langfuse_host " : " https://us.cloud.langfuse.com " ,
} ,
}
]
} ,
tpm_limit = None ,
rpm_limit = None ,
budget_duration = None ,
budget_reset_at = None ,
allowed_cache_controls = [ ] ,
permissions = { } ,
model_spend = { } ,
model_max_budget = { } ,
soft_budget_cooldown = False ,
litellm_budget_table = None ,
org_id = None ,
team_spend = None ,
team_alias = None ,
team_tpm_limit = None ,
team_rpm_limit = None ,
team_max_budget = None ,
team_models = [ ] ,
team_blocked = False ,
soft_budget = None ,
team_model_aliases = None ,
team_member_spend = None ,
team_metadata = None ,
end_user_id = None ,
end_user_tpm_limit = None ,
end_user_rpm_limit = None ,
end_user_max_budget = None ,
last_refreshed_at = None ,
api_key = None ,
user_role = None ,
allowed_model_region = None ,
parent_otel_span = None ,
) ,
" proxy_config " : proxy_config ,
" general_settings " : { } ,
" version " : " 0.0.0 " ,
}
new_data = await add_litellm_data_to_request ( * * data )
2024-09-10 07:44:37 +08:00
print ( " NEW DATA: {} " . format ( new_data ) )
2024-08-14 12:27:59 +08:00
assert " langfuse_public_key " in new_data
2024-08-22 04:36:33 +08:00
assert new_data [ " langfuse_public_key " ] == " my-mock-public-key "
2024-08-14 12:27:59 +08:00
assert " langfuse_secret_key " in new_data
2024-08-22 04:36:33 +08:00
assert new_data [ " langfuse_secret_key " ] == " my-mock-secret-key "
2024-08-18 01:46:59 +08:00
2024-09-10 07:44:37 +08:00
if expected_success_callbacks :
assert " success_callback " in new_data
assert new_data [ " success_callback " ] == expected_success_callbacks
if expected_failure_callbacks :
assert " failure_callback " in new_data
assert new_data [ " failure_callback " ] == expected_failure_callbacks
2024-08-18 01:46:59 +08:00
@pytest.mark.asyncio
async def test_gemini_pass_through_endpoint ( ) :
from starlette . datastructures import URL
from litellm . proxy . vertex_ai_endpoints . google_ai_studio_endpoints import (
Request ,
Response ,
gemini_proxy_route ,
)
body = b """
{
" contents " : [ {
" parts " : [ {
" text " : " The quick brown fox jumps over the lazy dog. "
} ]
} ]
}
"""
# Construct the scope dictionary
scope = {
" type " : " http " ,
" method " : " POST " ,
" path " : " /gemini/v1beta/models/gemini-1.5-flash:countTokens " ,
" query_string " : b " key=sk-1234 " ,
" headers " : [
( b " content-type " , b " application/json " ) ,
] ,
}
# Create a new Request object
async def async_receive ( ) :
return { " type " : " http.request " , " body " : body , " more_body " : False }
request = Request (
scope = scope ,
receive = async_receive ,
)
resp = await gemini_proxy_route (
endpoint = " v1beta/models/gemini-1.5-flash:countTokens?key=sk-1234 " ,
request = request ,
fastapi_response = Response ( ) ,
)
print ( resp . body )
2024-09-13 14:04:06 +08:00
2024-09-15 01:02:55 +08:00
@pytest.mark.parametrize ( " hidden " , [ True , False ] )
2024-09-13 14:04:06 +08:00
@pytest.mark.asyncio
2024-09-15 01:02:55 +08:00
async def test_proxy_model_group_alias_checks ( prisma_client , hidden ) :
2024-09-13 14:04:06 +08:00
"""
Check if model group alias is returned on
` / v1 / models `
` / v1 / model / info `
` / v1 / model_group / info `
"""
import json
from fastapi import HTTPException , Request , Response
from starlette . datastructures import URL
from litellm . proxy . proxy_server import model_group_info , model_info_v1 , model_list
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
proxy_config = getattr ( litellm . proxy . proxy_server , " proxy_config " )
_model_list = [
{
" model_name " : " gpt-3.5-turbo " ,
" litellm_params " : { " model " : " gpt-3.5-turbo " } ,
}
]
model_alias = " gpt-4 "
router = litellm . Router (
model_list = _model_list ,
2024-09-15 01:02:55 +08:00
model_group_alias = { model_alias : { " model " : " gpt-3.5-turbo " , " hidden " : hidden } } ,
2024-09-13 14:04:06 +08:00
)
setattr ( litellm . proxy . proxy_server , " llm_router " , router )
setattr ( litellm . proxy . proxy_server , " llm_model_list " , _model_list )
request = Request ( scope = { " type " : " http " , " method " : " POST " , " headers " : { } } )
request . _url = URL ( url = " /v1/models " )
resp = await model_list (
user_api_key_dict = UserAPIKeyAuth ( models = [ ] ) ,
)
2024-09-15 01:02:55 +08:00
if hidden :
assert len ( resp [ " data " ] ) == 1
else :
assert len ( resp [ " data " ] ) == 2
2024-09-13 14:04:06 +08:00
print ( resp )
resp = await model_info_v1 (
user_api_key_dict = UserAPIKeyAuth ( models = [ ] ) ,
)
models = resp [ " data " ]
is_model_alias_in_list = False
for item in models :
if model_alias == item [ " model_name " ] :
is_model_alias_in_list = True
2024-09-15 01:02:55 +08:00
if hidden :
assert is_model_alias_in_list is False
else :
assert is_model_alias_in_list
2024-09-13 14:04:06 +08:00
resp = await model_group_info (
user_api_key_dict = UserAPIKeyAuth ( models = [ ] ) ,
)
models = resp [ " data " ]
is_model_alias_in_list = False
for item in models :
if model_alias == item . model_group :
is_model_alias_in_list = True
2024-09-15 01:02:55 +08:00
if hidden :
assert is_model_alias_in_list is False
else :
assert is_model_alias_in_list , f " models: { models } "
2024-09-29 01:54:43 +08:00
@pytest.mark.asyncio
async def test_proxy_model_group_info_rerank ( prisma_client ) :
"""
Check if rerank model is returned on the following endpoints
` / v1 / models `
` / v1 / model / info `
` / v1 / model_group / info `
"""
import json
from fastapi import HTTPException , Request , Response
from starlette . datastructures import URL
from litellm . proxy . proxy_server import model_group_info , model_info_v1 , model_list
setattr ( litellm . proxy . proxy_server , " prisma_client " , prisma_client )
setattr ( litellm . proxy . proxy_server , " master_key " , " sk-1234 " )
await litellm . proxy . proxy_server . prisma_client . connect ( )
proxy_config = getattr ( litellm . proxy . proxy_server , " proxy_config " )
_model_list = [
{
" model_name " : " rerank-english-v3.0 " ,
" litellm_params " : { " model " : " cohere/rerank-english-v3.0 " } ,
" model_info " : {
" mode " : " rerank " ,
} ,
}
]
router = litellm . Router ( model_list = _model_list )
setattr ( litellm . proxy . proxy_server , " llm_router " , router )
setattr ( litellm . proxy . proxy_server , " llm_model_list " , _model_list )
request = Request ( scope = { " type " : " http " , " method " : " POST " , " headers " : { } } )
request . _url = URL ( url = " /v1/models " )
resp = await model_list (
user_api_key_dict = UserAPIKeyAuth ( models = [ ] ) ,
)
assert len ( resp [ " data " ] ) == 1
print ( resp )
resp = await model_info_v1 (
user_api_key_dict = UserAPIKeyAuth ( models = [ ] ) ,
)
models = resp [ " data " ]
assert models [ 0 ] [ " model_info " ] [ " mode " ] == " rerank "
resp = await model_group_info (
user_api_key_dict = UserAPIKeyAuth ( models = [ ] ) ,
)
print ( resp )
models = resp [ " data " ]
assert models [ 0 ] . mode == " rerank "