mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge branch 'main' into litellm_allow_setting_guardrails_config
This commit is contained in:
commit
b1e6cee000
12 changed files with 189 additions and 32 deletions
|
|
@ -113,6 +113,8 @@ ssl_verify: bool = True
|
|||
ssl_certificate: Optional[str] = None
|
||||
disable_streaming_logging: bool = False
|
||||
in_memory_llm_clients_cache: dict = {}
|
||||
### DEFAULT AZURE API VERSION ###
|
||||
AZURE_DEFAULT_API_VERSION = "2024-02-01" # this is updated to the latest
|
||||
### GUARDRAILS ###
|
||||
llamaguard_model_name: Optional[str] = None
|
||||
openai_moderations_model_name: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -75,16 +75,16 @@ class ServiceLogging(CustomLogger):
|
|||
await self.prometheusServicesLogger.async_service_success_hook(
|
||||
payload=payload
|
||||
)
|
||||
elif callback == "otel":
|
||||
from litellm.proxy.proxy_server import open_telemetry_logger
|
||||
|
||||
from litellm.proxy.proxy_server import open_telemetry_logger
|
||||
|
||||
if parent_otel_span is not None and open_telemetry_logger is not None:
|
||||
await open_telemetry_logger.async_service_success_hook(
|
||||
payload=payload,
|
||||
parent_otel_span=parent_otel_span,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
if parent_otel_span is not None and open_telemetry_logger is not None:
|
||||
await open_telemetry_logger.async_service_success_hook(
|
||||
payload=payload,
|
||||
parent_otel_span=parent_otel_span,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
async def async_service_failure_hook(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -295,7 +295,15 @@ def handle_prediction_response_streaming(prediction_url, api_token, print_verbos
|
|||
response_data = response.json()
|
||||
status = response_data["status"]
|
||||
if "output" in response_data:
|
||||
output_string = "".join(response_data["output"])
|
||||
try:
|
||||
output_string = "".join(response_data["output"])
|
||||
except Exception as e:
|
||||
raise ReplicateError(
|
||||
status_code=422,
|
||||
message="Unable to parse response. Got={}".format(
|
||||
response_data["output"]
|
||||
),
|
||||
)
|
||||
new_output = output_string[len(previous_output) :]
|
||||
print_verbose(f"New chunk: {new_output}")
|
||||
yield {"output": new_output, "status": status}
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm.utils import ModelResponse, EmbeddingResponse, get_secret, Usage
|
|||
import sys
|
||||
from copy import deepcopy
|
||||
import httpx # type: ignore
|
||||
import io
|
||||
from .prompt_templates.factory import prompt_factory, custom_prompt
|
||||
|
||||
|
||||
|
|
@ -25,10 +26,6 @@ class SagemakerError(Exception):
|
|||
) # Call the base class constructor with the parameters it needs
|
||||
|
||||
|
||||
import io
|
||||
import json
|
||||
|
||||
|
||||
class TokenIterator:
|
||||
def __init__(self, stream, acompletion: bool = False):
|
||||
if acompletion == False:
|
||||
|
|
@ -185,7 +182,8 @@ def completion(
|
|||
# I assume majority of users use .env for auth
|
||||
region_name = (
|
||||
get_secret("AWS_REGION_NAME")
|
||||
or "us-west-2" # default to us-west-2 if user not specified
|
||||
or aws_region_name # get region from config file if specified
|
||||
or "us-west-2" # default to us-west-2 if region not specified
|
||||
)
|
||||
client = boto3.client(
|
||||
service_name="sagemaker-runtime",
|
||||
|
|
@ -439,7 +437,8 @@ async def async_streaming(
|
|||
# I assume majority of users use .env for auth
|
||||
region_name = (
|
||||
get_secret("AWS_REGION_NAME")
|
||||
or "us-west-2" # default to us-west-2 if user not specified
|
||||
or aws_region_name # get region from config file if specified
|
||||
or "us-west-2" # default to us-west-2 if region not specified
|
||||
)
|
||||
_client = session.client(
|
||||
service_name="sagemaker-runtime",
|
||||
|
|
@ -506,7 +505,8 @@ async def async_completion(
|
|||
# I assume majority of users use .env for auth
|
||||
region_name = (
|
||||
get_secret("AWS_REGION_NAME")
|
||||
or "us-west-2" # default to us-west-2 if user not specified
|
||||
or aws_region_name # get region from config file if specified
|
||||
or "us-west-2" # default to us-west-2 if region not specified
|
||||
)
|
||||
_client = session.client(
|
||||
service_name="sagemaker-runtime",
|
||||
|
|
@ -661,7 +661,8 @@ def embedding(
|
|||
# I assume majority of users use .env for auth
|
||||
region_name = (
|
||||
get_secret("AWS_REGION_NAME")
|
||||
or "us-west-2" # default to us-west-2 if user not specified
|
||||
or aws_region_name # get region from config file if specified
|
||||
or "us-west-2" # default to us-west-2 if region not specified
|
||||
)
|
||||
client = boto3.client(
|
||||
service_name="sagemaker-runtime",
|
||||
|
|
|
|||
|
|
@ -4,6 +4,10 @@ model_list:
|
|||
model: "openai/*"
|
||||
mock_response: "Hello world!"
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["otel"]
|
||||
cache: True
|
||||
|
||||
general_settings:
|
||||
alerting: ["slack"]
|
||||
alerting_threshold: 10
|
||||
|
|
|
|||
|
|
@ -2,10 +2,10 @@ model_list:
|
|||
- model_name: claude-3-5-sonnet
|
||||
litellm_params:
|
||||
model: anthropic/claude-3-5-sonnet
|
||||
- model_name: gemini-1.5-flash-gemini
|
||||
litellm_params:
|
||||
model: vertex_ai_beta/gemini-1.5-flash
|
||||
api_base: https://gateway.ai.cloudflare.com/v1/fa4cdcab1f32b95ca3b53fd36043d691/test/google-vertex-ai/v1/projects/adroit-crow-413218/locations/us-central1/publishers/google/models/gemini-1.5-flash
|
||||
# - model_name: gemini-1.5-flash-gemini
|
||||
# litellm_params:
|
||||
# model: vertex_ai_beta/gemini-1.5-flash
|
||||
# api_base: https://gateway.ai.cloudflare.com/v1/fa4cdcab1f32b95ca3b53fd36043d691/test/google-vertex-ai/v1/projects/adroit-crow-413218/locations/us-central1/publishers/google/models/gemini-1.5-flash
|
||||
- litellm_params:
|
||||
api_base: http://0.0.0.0:8080
|
||||
api_key: ''
|
||||
|
|
|
|||
|
|
@ -218,6 +218,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/v2/model/info",
|
||||
"/v2/key/info",
|
||||
"/model_group/info",
|
||||
"/health",
|
||||
]
|
||||
|
||||
# NOTE: ROUTES ONLY FOR MASTER KEY - only the Master Key should be able to Reset Spend
|
||||
|
|
|
|||
|
|
@ -3437,7 +3437,7 @@ class Router:
|
|||
if azure_ad_token.startswith("oidc/"):
|
||||
azure_ad_token = get_azure_ad_token_from_oidc(azure_ad_token)
|
||||
if api_version is None:
|
||||
api_version = "2023-07-01-preview"
|
||||
api_version = litellm.AZURE_DEFAULT_API_VERSION
|
||||
|
||||
if "gateway.ai.cloudflare.com" in api_base:
|
||||
if not api_base.endswith("/"):
|
||||
|
|
|
|||
|
|
@ -1,6 +1,9 @@
|
|||
import sys, os, uuid
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import traceback
|
||||
import uuid
|
||||
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
|
@ -9,12 +12,15 @@ 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, aembedding
|
||||
from litellm.caching import Cache
|
||||
import asyncio
|
||||
import hashlib
|
||||
import random
|
||||
import hashlib, asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm import aembedding, completion, embedding
|
||||
from litellm.caching import Cache
|
||||
|
||||
# litellm.set_verbose=True
|
||||
|
||||
|
|
@ -656,6 +662,7 @@ def test_redis_cache_completion():
|
|||
assert response1.created == response2.created
|
||||
assert response1.choices[0].message.content == response2.choices[0].message.content
|
||||
|
||||
|
||||
# test_redis_cache_completion()
|
||||
|
||||
|
||||
|
|
@ -877,6 +884,7 @@ async def test_redis_cache_acompletion_stream_bedrock():
|
|||
print(e)
|
||||
raise e
|
||||
|
||||
|
||||
def test_disk_cache_completion():
|
||||
litellm.set_verbose = False
|
||||
|
||||
|
|
@ -925,7 +933,7 @@ def test_disk_cache_completion():
|
|||
litellm.success_callback = []
|
||||
litellm._async_success_callback = []
|
||||
|
||||
# 1 & 2 should be exactly the same
|
||||
# 1 & 2 should be exactly the same
|
||||
# 1 & 3 should be different, since input params are diff
|
||||
if (
|
||||
response1["choices"][0]["message"]["content"]
|
||||
|
|
@ -1569,3 +1577,37 @@ async def test_redis_semantic_cache_acompletion():
|
|||
)
|
||||
print(f"response2: {response2}")
|
||||
assert response1.id == response2.id
|
||||
|
||||
|
||||
def test_caching_redis_simple(caplog):
|
||||
"""
|
||||
Relevant issue - https://github.com/BerriAI/litellm/issues/4511
|
||||
"""
|
||||
litellm.cache = Cache(
|
||||
type="redis", url=os.getenv("REDIS_SSL_URL")
|
||||
) # passing `supported_call_types = ["completion"]` has no effect
|
||||
|
||||
s = time.time()
|
||||
x = completion(
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hello, how are you? Wink"}],
|
||||
stream=True,
|
||||
)
|
||||
for m in x:
|
||||
print(m)
|
||||
print(time.time() - s)
|
||||
|
||||
s2 = time.time()
|
||||
x = completion(
|
||||
model="gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hello, how are you? Wink"}],
|
||||
stream=True,
|
||||
)
|
||||
for m in x:
|
||||
print(m)
|
||||
print(time.time() - s2)
|
||||
|
||||
captured_logs = [rec.message for rec in caplog.records]
|
||||
|
||||
assert "LiteLLM Redis Caching: async set" not in captured_logs
|
||||
assert "ServiceLogging.async_service_success_hook" not in captured_logs
|
||||
|
|
|
|||
|
|
@ -512,6 +512,106 @@ def sagemaker_test_completion():
|
|||
|
||||
# sagemaker_test_completion()
|
||||
|
||||
|
||||
def test_sagemaker_default_region(mocker):
|
||||
"""
|
||||
If no regions are specified in config or in environment, the default region is us-west-2
|
||||
"""
|
||||
mock_client = mocker.patch("boto3.client")
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="sagemaker/mock-endpoint",
|
||||
messages=[
|
||||
{
|
||||
"content": "Hello, world!",
|
||||
"role": "user"
|
||||
}
|
||||
]
|
||||
)
|
||||
except Exception:
|
||||
pass # expected serialization exception because AWS client was replaced with a Mock
|
||||
assert mock_client.call_args.kwargs["region_name"] == "us-west-2"
|
||||
|
||||
# test_sagemaker_default_region()
|
||||
|
||||
|
||||
def test_sagemaker_environment_region(mocker):
|
||||
"""
|
||||
If a region is specified in the environment, use that region instead of us-west-2
|
||||
"""
|
||||
expected_region = "us-east-1"
|
||||
os.environ["AWS_REGION_NAME"] = expected_region
|
||||
mock_client = mocker.patch("boto3.client")
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="sagemaker/mock-endpoint",
|
||||
messages=[
|
||||
{
|
||||
"content": "Hello, world!",
|
||||
"role": "user"
|
||||
}
|
||||
]
|
||||
)
|
||||
except Exception:
|
||||
pass # expected serialization exception because AWS client was replaced with a Mock
|
||||
del os.environ["AWS_REGION_NAME"] # cleanup
|
||||
assert mock_client.call_args.kwargs["region_name"] == expected_region
|
||||
|
||||
# test_sagemaker_environment_region()
|
||||
|
||||
|
||||
def test_sagemaker_config_region(mocker):
|
||||
"""
|
||||
If a region is specified as part of the optional parameters of the completion, including as
|
||||
part of the config file, then use that region instead of us-west-2
|
||||
"""
|
||||
expected_region = "us-east-1"
|
||||
mock_client = mocker.patch("boto3.client")
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="sagemaker/mock-endpoint",
|
||||
messages=[
|
||||
{
|
||||
"content": "Hello, world!",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
aws_region_name=expected_region,
|
||||
)
|
||||
except Exception:
|
||||
pass # expected serialization exception because AWS client was replaced with a Mock
|
||||
assert mock_client.call_args.kwargs["region_name"] == expected_region
|
||||
|
||||
# test_sagemaker_config_region()
|
||||
|
||||
|
||||
def test_sagemaker_config_and_environment_region(mocker):
|
||||
"""
|
||||
If both the environment and config file specify a region, the environment region is expected
|
||||
"""
|
||||
expected_region = "us-east-1"
|
||||
unexpected_region = "us-east-2"
|
||||
os.environ["AWS_REGION_NAME"] = expected_region
|
||||
mock_client = mocker.patch("boto3.client")
|
||||
try:
|
||||
response = litellm.completion(
|
||||
model="sagemaker/mock-endpoint",
|
||||
messages=[
|
||||
{
|
||||
"content": "Hello, world!",
|
||||
"role": "user"
|
||||
}
|
||||
],
|
||||
aws_region_name=unexpected_region,
|
||||
)
|
||||
except Exception:
|
||||
pass # expected serialization exception because AWS client was replaced with a Mock
|
||||
del os.environ["AWS_REGION_NAME"] # cleanup
|
||||
assert mock_client.call_args.kwargs["region_name"] == expected_region
|
||||
|
||||
# test_sagemaker_config_and_environment_region()
|
||||
|
||||
|
||||
# Bedrock
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1415,7 +1415,6 @@ def test_bedrock_claude_3_streaming():
|
|||
"gpt-3.5-turbo",
|
||||
"databricks/databricks-dbrx-instruct", # databricks
|
||||
"predibase/llama-3-8b-instruct", # predibase
|
||||
"replicate/meta/meta-llama-3-8b-instruct", # replicate
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -3634,7 +3634,7 @@ def get_model_region(
|
|||
model=_model,
|
||||
api_key=litellm_params.api_key,
|
||||
api_base=litellm_params.api_base,
|
||||
api_version=litellm_params.api_version or "2023-07-01-preview",
|
||||
api_version=litellm_params.api_version or litellm.AZURE_DEFAULT_API_VERSION,
|
||||
timeout=10,
|
||||
mode=mode or "chat",
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue