mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: req changes with relevant test cases
This commit is contained in:
parent
3761598a02
commit
8bdbda0d2c
2 changed files with 444 additions and 149 deletions
|
|
@ -84,7 +84,9 @@ def _get_redis_cluster_kwargs(client=None):
|
|||
available_args.append("ssl_cert_reqs")
|
||||
available_args.append("ssl_check_hostname")
|
||||
available_args.append("ssl_ca_certs")
|
||||
available_args.append("redis_connect_func") # Needed for sync clusters and IAM detection
|
||||
available_args.append(
|
||||
"redis_connect_func"
|
||||
) # Needed for sync clusters and IAM detection
|
||||
available_args.append("gcp_service_account")
|
||||
available_args.append("gcp_ssl_ca_certs")
|
||||
available_args.append("azure_redis_ad_token")
|
||||
|
|
@ -301,7 +303,9 @@ def get_redis_url_from_environment():
|
|||
return os.environ["REDIS_URL"]
|
||||
|
||||
if "REDIS_HOST" not in os.environ or "REDIS_PORT" not in os.environ:
|
||||
raise ValueError("Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis.")
|
||||
raise ValueError(
|
||||
"Either 'REDIS_URL' or both 'REDIS_HOST' and 'REDIS_PORT' must be specified for Redis."
|
||||
)
|
||||
|
||||
if "REDIS_SSL" in os.environ and os.environ["REDIS_SSL"].lower() == "true":
|
||||
redis_protocol = "rediss"
|
||||
|
|
@ -348,9 +352,9 @@ def _get_redis_client_logic(**env_overrides): # noqa: PLR0915
|
|||
if _sentinel_nodes is not None and isinstance(_sentinel_nodes, str):
|
||||
redis_kwargs["sentinel_nodes"] = json.loads(_sentinel_nodes)
|
||||
|
||||
_sentinel_password: Optional[str] = redis_kwargs.get("sentinel_password", None) or get_secret_str(
|
||||
"REDIS_SENTINEL_PASSWORD"
|
||||
)
|
||||
_sentinel_password: Optional[str] = redis_kwargs.get(
|
||||
"sentinel_password", None
|
||||
) or get_secret_str("REDIS_SENTINEL_PASSWORD")
|
||||
|
||||
if _sentinel_password is not None:
|
||||
redis_kwargs["sentinel_password"] = _sentinel_password
|
||||
|
|
@ -363,11 +367,17 @@ def _get_redis_client_logic(**env_overrides): # noqa: PLR0915
|
|||
redis_kwargs["service_name"] = _service_name
|
||||
|
||||
# Handle GCP IAM authentication
|
||||
_gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
|
||||
_gcp_ssl_ca_certs = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str("REDIS_GCP_SSL_CA_CERTS")
|
||||
_gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str(
|
||||
"REDIS_GCP_SERVICE_ACCOUNT"
|
||||
)
|
||||
_gcp_ssl_ca_certs = redis_kwargs.get("gcp_ssl_ca_certs") or get_secret_str(
|
||||
"REDIS_GCP_SSL_CA_CERTS"
|
||||
)
|
||||
|
||||
if _gcp_service_account is not None:
|
||||
verbose_logger.debug("Setting up GCP IAM authentication for Redis with service account.")
|
||||
verbose_logger.debug(
|
||||
"Setting up GCP IAM authentication for Redis with service account."
|
||||
)
|
||||
redis_kwargs["redis_connect_func"] = create_gcp_iam_redis_connect_func(
|
||||
service_account=_gcp_service_account, ssl_ca_certs=_gcp_ssl_ca_certs
|
||||
)
|
||||
|
|
@ -383,12 +393,37 @@ def _get_redis_client_logic(**env_overrides): # noqa: PLR0915
|
|||
redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs
|
||||
|
||||
# Handle Azure AD authentication (after GCP IAM block)
|
||||
_azure_redis_ad_token = redis_kwargs.get("azure_redis_ad_token") or get_secret("REDIS_AZURE_AD_TOKEN")
|
||||
_azure_redis_ad_token = redis_kwargs.get("azure_redis_ad_token") or get_secret(
|
||||
"REDIS_AZURE_AD_TOKEN"
|
||||
)
|
||||
|
||||
if _azure_redis_ad_token is not None and str(_azure_redis_ad_token).lower() == "true":
|
||||
_azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str("AZURE_CLIENT_ID")
|
||||
_azure_tenant_id = redis_kwargs.get("azure_tenant_id") or get_secret_str("AZURE_TENANT_ID")
|
||||
_azure_client_secret = redis_kwargs.get("azure_client_secret") or get_secret_str("AZURE_CLIENT_SECRET")
|
||||
if (
|
||||
_azure_redis_ad_token is not None
|
||||
and str(_azure_redis_ad_token).lower() == "true"
|
||||
and _gcp_service_account is not None
|
||||
):
|
||||
verbose_logger.warning(
|
||||
"Both GCP IAM (gcp_service_account) and Azure AD (azure_redis_ad_token) are configured for Redis. "
|
||||
"Using GCP IAM. Remove one to avoid misconfiguration."
|
||||
)
|
||||
# Clean up Azure-specific kwargs even though we're not using Azure AD
|
||||
redis_kwargs.pop("azure_redis_ad_token", None)
|
||||
redis_kwargs.pop("azure_client_id", None)
|
||||
redis_kwargs.pop("azure_tenant_id", None)
|
||||
redis_kwargs.pop("azure_client_secret", None)
|
||||
elif (
|
||||
_azure_redis_ad_token is not None
|
||||
and str(_azure_redis_ad_token).lower() == "true"
|
||||
):
|
||||
_azure_client_id = redis_kwargs.get("azure_client_id") or get_secret_str(
|
||||
"AZURE_CLIENT_ID"
|
||||
)
|
||||
_azure_tenant_id = redis_kwargs.get("azure_tenant_id") or get_secret_str(
|
||||
"AZURE_TENANT_ID"
|
||||
)
|
||||
_azure_client_secret = redis_kwargs.get(
|
||||
"azure_client_secret"
|
||||
) or get_secret_str("AZURE_CLIENT_SECRET")
|
||||
|
||||
verbose_logger.debug("Setting up Azure AD authentication for Redis.")
|
||||
redis_kwargs["redis_connect_func"] = create_azure_ad_redis_connect_func(
|
||||
|
|
@ -415,7 +450,9 @@ def _get_redis_client_logic(**env_overrides): # noqa: PLR0915
|
|||
redis_kwargs.pop("password", None)
|
||||
elif "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None:
|
||||
pass
|
||||
elif "sentinel_nodes" in redis_kwargs and redis_kwargs["sentinel_nodes"] is not None:
|
||||
elif (
|
||||
"sentinel_nodes" in redis_kwargs and redis_kwargs["sentinel_nodes"] is not None
|
||||
):
|
||||
pass
|
||||
elif "host" not in redis_kwargs or redis_kwargs["host"] is None:
|
||||
raise ValueError("Either 'host' or 'url' must be specified for redis.")
|
||||
|
|
@ -458,7 +495,9 @@ def _init_redis_sentinel(redis_kwargs) -> redis.Redis:
|
|||
service_name = redis_kwargs.get("service_name")
|
||||
|
||||
if not sentinel_nodes or not service_name:
|
||||
raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.")
|
||||
raise ValueError(
|
||||
"Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel."
|
||||
)
|
||||
|
||||
verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.")
|
||||
|
||||
|
|
@ -480,7 +519,9 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis:
|
|||
service_name = redis_kwargs.get("service_name")
|
||||
|
||||
if not sentinel_nodes or not service_name:
|
||||
raise ValueError("Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel.")
|
||||
raise ValueError(
|
||||
"Both 'sentinel_nodes' and 'service_name' are required for Redis Sentinel."
|
||||
)
|
||||
|
||||
verbose_logger.debug("init_redis_sentinel: sentinel nodes are being initialized.")
|
||||
|
||||
|
|
@ -532,7 +573,9 @@ def get_redis_async_client( # noqa: PLR0915
|
|||
url_kwargs[arg] = redis_kwargs[arg]
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
"REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(arg)
|
||||
"REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format(
|
||||
arg
|
||||
)
|
||||
)
|
||||
return async_redis.Redis.from_url(**url_kwargs)
|
||||
|
||||
|
|
@ -554,7 +597,9 @@ def get_redis_async_client( # noqa: PLR0915
|
|||
if redis_connect_func and hasattr(redis_connect_func, "_gcp_service_account"):
|
||||
gcp_service_account = redis_connect_func._gcp_service_account
|
||||
else:
|
||||
gcp_service_account = redis_kwargs.get("gcp_service_account") or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
|
||||
gcp_service_account = redis_kwargs.get(
|
||||
"gcp_service_account"
|
||||
) or get_secret_str("REDIS_GCP_SERVICE_ACCOUNT")
|
||||
|
||||
verbose_logger.debug(
|
||||
f"DEBUG: Redis cluster kwargs: redis_connect_func={redis_connect_func is not None}, gcp_service_account_provided={gcp_service_account is not None}"
|
||||
|
|
@ -569,17 +614,23 @@ def get_redis_async_client( # noqa: PLR0915
|
|||
# Generate IAM access token using the helper function
|
||||
access_token = _generate_gcp_iam_access_token(gcp_service_account)
|
||||
cluster_kwargs["password"] = access_token
|
||||
verbose_logger.debug("DEBUG: Successfully generated GCP IAM access token for async Redis cluster")
|
||||
verbose_logger.debug(
|
||||
"DEBUG: Successfully generated GCP IAM access token for async Redis cluster"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Failed to generate GCP IAM access token: {e}")
|
||||
from redis.exceptions import AuthenticationError
|
||||
|
||||
raise AuthenticationError("Failed to generate GCP IAM access token")
|
||||
# Handle Azure AD authentication for async clusters
|
||||
elif redis_connect_func and hasattr(redis_connect_func, "_azure_redis_ad_token"):
|
||||
elif redis_connect_func and hasattr(
|
||||
redis_connect_func, "_azure_redis_ad_token"
|
||||
):
|
||||
_az_client_id = getattr(redis_connect_func, "_azure_client_id", None)
|
||||
_az_tenant_id = getattr(redis_connect_func, "_azure_tenant_id", None)
|
||||
_az_client_secret = getattr(redis_connect_func, "_azure_client_secret", None)
|
||||
_az_client_secret = getattr(
|
||||
redis_connect_func, "_azure_client_secret", None
|
||||
)
|
||||
|
||||
verbose_logger.debug("Generating Azure AD token for async Redis cluster")
|
||||
try:
|
||||
|
|
@ -593,12 +644,16 @@ def get_redis_async_client( # noqa: PLR0915
|
|||
_username = os.environ.get("REDIS_USERNAME", "")
|
||||
if _username:
|
||||
cluster_kwargs["username"] = _username
|
||||
verbose_logger.debug("Successfully generated Azure AD token for async Redis cluster")
|
||||
verbose_logger.debug(
|
||||
"Successfully generated Azure AD token for async Redis cluster"
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Failed to generate Azure AD access token: {e}")
|
||||
from redis.exceptions import AuthenticationError
|
||||
|
||||
raise AuthenticationError("Failed to generate Azure AD access token for Redis")
|
||||
raise AuthenticationError(
|
||||
"Failed to generate Azure AD access token for Redis"
|
||||
)
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"DEBUG: Not using GCP/Azure AD IAM auth - redis_connect_func={redis_connect_func is not None}"
|
||||
|
|
@ -654,7 +709,9 @@ def get_redis_connection_pool(**env_overrides):
|
|||
redis_kwargs.pop("ssl", None)
|
||||
redis_kwargs["connection_class"] = connection_class
|
||||
redis_kwargs.pop("startup_nodes", None)
|
||||
return async_redis.BlockingConnectionPool(timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs)
|
||||
return async_redis.BlockingConnectionPool(
|
||||
timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs
|
||||
)
|
||||
|
||||
|
||||
def _pretty_print_redis_config(redis_kwargs: dict) -> None:
|
||||
|
|
|
|||
|
|
@ -6,7 +6,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||
import pytest
|
||||
from jsonschema import validate
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
from litellm.proxy.utils import is_valid_api_key
|
||||
|
|
@ -34,13 +36,28 @@ def test_check_provider_match_azure_ai_allows_openai_and_azure():
|
|||
This is needed for Azure Model Router which can route to OpenAI models.
|
||||
"""
|
||||
# azure_ai should match openai models
|
||||
assert _check_provider_match(model_info={"litellm_provider": "openai"}, custom_llm_provider="azure_ai") is True
|
||||
assert (
|
||||
_check_provider_match(
|
||||
model_info={"litellm_provider": "openai"}, custom_llm_provider="azure_ai"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
# azure_ai should match azure models
|
||||
assert _check_provider_match(model_info={"litellm_provider": "azure"}, custom_llm_provider="azure_ai") is True
|
||||
assert (
|
||||
_check_provider_match(
|
||||
model_info={"litellm_provider": "azure"}, custom_llm_provider="azure_ai"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
# azure_ai should NOT match other providers
|
||||
assert _check_provider_match(model_info={"litellm_provider": "anthropic"}, custom_llm_provider="azure_ai") is False
|
||||
assert (
|
||||
_check_provider_match(
|
||||
model_info={"litellm_provider": "anthropic"}, custom_llm_provider="azure_ai"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_check_provider_match_github_allows_upstream_provider_metadata():
|
||||
|
|
@ -75,11 +92,19 @@ def test_check_provider_match_github_allows_upstream_provider_metadata():
|
|||
|
||||
def test_supports_function_calling_github_openai_alias():
|
||||
assert litellm.utils.supports_function_calling(model="github/gpt-4o-mini") is True
|
||||
assert litellm.utils.supports_function_calling(model="gpt-4o-mini", custom_llm_provider="github") is True
|
||||
assert (
|
||||
litellm.utils.supports_function_calling(
|
||||
model="gpt-4o-mini", custom_llm_provider="github"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_supports_function_calling_github_anthropic_alias():
|
||||
assert litellm.utils.supports_function_calling(model="github/claude-3-5-sonnet-latest") is True
|
||||
assert (
|
||||
litellm.utils.supports_function_calling(model="github/claude-3-5-sonnet-latest")
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_supports_function_calling_deepinfra_llama():
|
||||
|
|
@ -96,7 +121,12 @@ def test_supports_function_calling_deepinfra_llama():
|
|||
|
||||
|
||||
def test_supports_function_calling_unknown_github_alias_returns_false():
|
||||
assert litellm.utils.supports_function_calling(model="github/non-existent-model-for-capability-check") is False
|
||||
assert (
|
||||
litellm.utils.supports_function_calling(
|
||||
model="github/non-existent-model-for-capability-check"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
||||
def test_get_optional_params_image_gen():
|
||||
|
|
@ -149,7 +179,9 @@ def test_get_optional_params_image_gen_vertex_ai_size():
|
|||
drop_params=True,
|
||||
)
|
||||
assert optional_params is not None
|
||||
assert "aspectRatio" not in optional_params # aspectRatio should not be set if size is not provided
|
||||
assert (
|
||||
"aspectRatio" not in optional_params
|
||||
) # aspectRatio should not be set if size is not provided
|
||||
assert optional_params["sampleCount"] == 1
|
||||
|
||||
|
||||
|
|
@ -170,19 +202,26 @@ def test_all_model_configs():
|
|||
VertexAILlama3Config,
|
||||
)
|
||||
|
||||
assert "max_completion_tokens" in VertexAILlama3Config().get_supported_openai_params(model="llama3")
|
||||
assert VertexAILlama3Config().map_openai_params({"max_completion_tokens": 10}, {}, "llama3", drop_params=False) == {
|
||||
"max_tokens": 10
|
||||
}
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in VertexAILlama3Config().get_supported_openai_params(model="llama3")
|
||||
)
|
||||
assert VertexAILlama3Config().map_openai_params(
|
||||
{"max_completion_tokens": 10}, {}, "llama3", drop_params=False
|
||||
) == {"max_tokens": 10}
|
||||
|
||||
assert "max_completion_tokens" in VertexAIAi21Config().get_supported_openai_params(model="jamba-1.5-mini@001")
|
||||
assert "max_completion_tokens" in VertexAIAi21Config().get_supported_openai_params(
|
||||
model="jamba-1.5-mini@001"
|
||||
)
|
||||
assert VertexAIAi21Config().map_openai_params(
|
||||
{"max_completion_tokens": 10}, {}, "jamba-1.5-mini@001", drop_params=False
|
||||
) == {"max_tokens": 10}
|
||||
|
||||
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
|
||||
|
||||
assert "max_completion_tokens" in FireworksAIConfig().get_supported_openai_params(model="llama3")
|
||||
assert "max_completion_tokens" in FireworksAIConfig().get_supported_openai_params(
|
||||
model="llama3"
|
||||
)
|
||||
assert FireworksAIConfig().map_openai_params(
|
||||
model="llama3",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -192,7 +231,9 @@ def test_all_model_configs():
|
|||
|
||||
from litellm.llms.nvidia_nim.chat.transformation import NvidiaNimConfig
|
||||
|
||||
assert "max_completion_tokens" in NvidiaNimConfig().get_supported_openai_params(model="llama3")
|
||||
assert "max_completion_tokens" in NvidiaNimConfig().get_supported_openai_params(
|
||||
model="llama3"
|
||||
)
|
||||
assert NvidiaNimConfig().map_openai_params(
|
||||
model="llama3",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -202,7 +243,9 @@ def test_all_model_configs():
|
|||
|
||||
from litellm.llms.ollama.chat.transformation import OllamaChatConfig
|
||||
|
||||
assert "max_completion_tokens" in OllamaChatConfig().get_supported_openai_params(model="llama3")
|
||||
assert "max_completion_tokens" in OllamaChatConfig().get_supported_openai_params(
|
||||
model="llama3"
|
||||
)
|
||||
assert OllamaChatConfig().map_openai_params(
|
||||
model="llama3",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -212,7 +255,9 @@ def test_all_model_configs():
|
|||
|
||||
from litellm.llms.predibase.chat.transformation import PredibaseConfig
|
||||
|
||||
assert "max_completion_tokens" in PredibaseConfig().get_supported_openai_params(model="llama3")
|
||||
assert "max_completion_tokens" in PredibaseConfig().get_supported_openai_params(
|
||||
model="llama3"
|
||||
)
|
||||
assert PredibaseConfig().map_openai_params(
|
||||
model="llama3",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -224,7 +269,10 @@ def test_all_model_configs():
|
|||
CodestralTextCompletionConfig,
|
||||
)
|
||||
|
||||
assert "max_completion_tokens" in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3")
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3")
|
||||
)
|
||||
assert CodestralTextCompletionConfig().map_openai_params(
|
||||
model="llama3",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -236,7 +284,9 @@ def test_all_model_configs():
|
|||
VolcEngineChatConfig as VolcEngineConfig,
|
||||
)
|
||||
|
||||
assert "max_completion_tokens" in VolcEngineConfig().get_supported_openai_params(model="llama3")
|
||||
assert "max_completion_tokens" in VolcEngineConfig().get_supported_openai_params(
|
||||
model="llama3"
|
||||
)
|
||||
assert VolcEngineConfig().map_openai_params(
|
||||
model="llama3",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -246,7 +296,9 @@ def test_all_model_configs():
|
|||
|
||||
from litellm.llms.ai21.chat.transformation import AI21ChatConfig
|
||||
|
||||
assert "max_completion_tokens" in AI21ChatConfig().get_supported_openai_params("jamba-1.5-mini@001")
|
||||
assert "max_completion_tokens" in AI21ChatConfig().get_supported_openai_params(
|
||||
"jamba-1.5-mini@001"
|
||||
)
|
||||
assert AI21ChatConfig().map_openai_params(
|
||||
model="jamba-1.5-mini@001",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -256,7 +308,9 @@ def test_all_model_configs():
|
|||
|
||||
from litellm.llms.azure.chat.gpt_transformation import AzureOpenAIConfig
|
||||
|
||||
assert "max_completion_tokens" in AzureOpenAIConfig().get_supported_openai_params(model="gpt-3.5-turbo")
|
||||
assert "max_completion_tokens" in AzureOpenAIConfig().get_supported_openai_params(
|
||||
model="gpt-3.5-turbo"
|
||||
)
|
||||
assert AzureOpenAIConfig().map_openai_params(
|
||||
model="gpt-3.5-turbo",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -267,8 +321,11 @@ def test_all_model_configs():
|
|||
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
|
||||
assert "max_completion_tokens" in AmazonConverseConfig().get_supported_openai_params(
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in AmazonConverseConfig().get_supported_openai_params(
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
)
|
||||
)
|
||||
assert AmazonConverseConfig().map_openai_params(
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
|
|
@ -281,7 +338,10 @@ def test_all_model_configs():
|
|||
CodestralTextCompletionConfig,
|
||||
)
|
||||
|
||||
assert "max_completion_tokens" in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3")
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3")
|
||||
)
|
||||
assert CodestralTextCompletionConfig().map_openai_params(
|
||||
model="llama3",
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -291,8 +351,11 @@ def test_all_model_configs():
|
|||
|
||||
from litellm import AmazonAnthropicClaudeConfig, AmazonAnthropicConfig
|
||||
|
||||
assert "max_completion_tokens" in AmazonAnthropicClaudeConfig().get_supported_openai_params(
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in AmazonAnthropicClaudeConfig().get_supported_openai_params(
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
)
|
||||
)
|
||||
|
||||
assert AmazonAnthropicClaudeConfig().map_openai_params(
|
||||
|
|
@ -302,7 +365,10 @@ def test_all_model_configs():
|
|||
drop_params=False,
|
||||
) == {"max_tokens": 10}
|
||||
|
||||
assert "max_completion_tokens" in AmazonAnthropicConfig().get_supported_openai_params(model="")
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in AmazonAnthropicConfig().get_supported_openai_params(model="")
|
||||
)
|
||||
|
||||
assert AmazonAnthropicConfig().map_openai_params(
|
||||
non_default_params={"max_completion_tokens": 10},
|
||||
|
|
@ -326,8 +392,11 @@ def test_all_model_configs():
|
|||
VertexAIAnthropicConfig,
|
||||
)
|
||||
|
||||
assert "max_completion_tokens" in VertexAIAnthropicConfig().get_supported_openai_params(
|
||||
model="claude-3-5-sonnet-20240620"
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in VertexAIAnthropicConfig().get_supported_openai_params(
|
||||
model="claude-3-5-sonnet-20240620"
|
||||
)
|
||||
)
|
||||
|
||||
assert VertexAIAnthropicConfig().map_openai_params(
|
||||
|
|
@ -342,7 +411,9 @@ def test_all_model_configs():
|
|||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(model="gemini-1.0-pro")
|
||||
assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(
|
||||
model="gemini-1.0-pro"
|
||||
)
|
||||
|
||||
assert VertexGeminiConfig().map_openai_params(
|
||||
model="gemini-1.0-pro",
|
||||
|
|
@ -351,7 +422,12 @@ def test_all_model_configs():
|
|||
drop_params=False,
|
||||
) == {"max_output_tokens": 10}
|
||||
|
||||
assert "max_completion_tokens" in GoogleAIStudioGeminiConfig().get_supported_openai_params(model="gemini-1.0-pro")
|
||||
assert (
|
||||
"max_completion_tokens"
|
||||
in GoogleAIStudioGeminiConfig().get_supported_openai_params(
|
||||
model="gemini-1.0-pro"
|
||||
)
|
||||
)
|
||||
|
||||
assert GoogleAIStudioGeminiConfig().map_openai_params(
|
||||
model="gemini-1.0-pro",
|
||||
|
|
@ -360,7 +436,9 @@ def test_all_model_configs():
|
|||
drop_params=False,
|
||||
) == {"max_output_tokens": 10}
|
||||
|
||||
assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(model="gemini-1.0-pro")
|
||||
assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(
|
||||
model="gemini-1.0-pro"
|
||||
)
|
||||
|
||||
assert VertexGeminiConfig().map_openai_params(
|
||||
model="gemini-1.0-pro",
|
||||
|
|
@ -386,7 +464,9 @@ def test_anthropic_web_search_in_model_info():
|
|||
|
||||
model_info = get_model_info(model)
|
||||
assert model_info is not None
|
||||
assert model_info["supports_web_search"] is True, f"Model {model} should support web search"
|
||||
assert (
|
||||
model_info["supports_web_search"] is True
|
||||
), f"Model {model} should support web search"
|
||||
assert (
|
||||
model_info["search_context_cost_per_query"] is not None
|
||||
), f"Model {model} should have a search context cost per query"
|
||||
|
|
@ -491,7 +571,9 @@ def validate_model_cost_values(model_data, exceptions=None):
|
|||
continue
|
||||
|
||||
if isinstance(cost_value, (int, float)) and cost_value > 1:
|
||||
violations.append(f"Model '{model_id}' has {field} = {cost_value} which exceeds 1")
|
||||
violations.append(
|
||||
f"Model '{model_id}' has {field} = {cost_value} which exceeds 1"
|
||||
)
|
||||
|
||||
# Check nested cost fields
|
||||
for field in nested_cost_fields:
|
||||
|
|
@ -535,7 +617,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"cache_read_input_token_cost": {"type": "number"},
|
||||
"cache_read_input_token_cost_above_200k_tokens": {"type": "number"},
|
||||
"cache_read_input_token_cost_above_272k_tokens": {"type": "number"},
|
||||
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": {"type": "number"},
|
||||
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": {
|
||||
"type": "number"
|
||||
},
|
||||
"cache_read_input_audio_token_cost": {"type": "number"},
|
||||
"cache_read_input_token_cost_per_audio_token": {"type": "number"},
|
||||
"cache_read_input_image_token_cost": {"type": "number"},
|
||||
|
|
@ -553,8 +637,12 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"input_cost_per_token_above_272k_tokens": {"type": "number"},
|
||||
"cache_read_input_token_cost_flex": {"type": "number"},
|
||||
"cache_read_input_token_cost_priority": {"type": "number"},
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": {"type": "number"},
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": {"type": "number"},
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": {
|
||||
"type": "number"
|
||||
},
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority": {
|
||||
"type": "number"
|
||||
},
|
||||
"input_cost_per_token_flex": {"type": "number"},
|
||||
"input_cost_per_token_priority": {"type": "number"},
|
||||
"input_cost_per_token_above_200k_tokens_priority": {"type": "number"},
|
||||
|
|
@ -575,7 +663,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"input_cost_per_token_cache_hit": {"type": "number"},
|
||||
"input_cost_per_video_per_second": {"type": "number"},
|
||||
"input_cost_per_video_per_second_above_8s_interval": {"type": "number"},
|
||||
"input_cost_per_video_per_second_above_15s_interval": {"type": "number"},
|
||||
"input_cost_per_video_per_second_above_15s_interval": {
|
||||
"type": "number"
|
||||
},
|
||||
"input_cost_per_video_per_second_above_128k_tokens": {"type": "number"},
|
||||
"input_dbu_cost_per_token": {"type": "number"},
|
||||
"annotation_cost_per_page": {"type": "number"},
|
||||
|
|
@ -630,9 +720,13 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
"output_cost_per_token_above_200k_tokens": {"type": "number"},
|
||||
"output_cost_per_token_above_272k_tokens": {"type": "number"},
|
||||
"output_cost_per_image_above_1024_and_1024_pixels": {"type": "number"},
|
||||
"output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": {"type": "number"},
|
||||
"output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": {
|
||||
"type": "number"
|
||||
},
|
||||
"output_cost_per_image_above_512_and_512_pixels": {"type": "number"},
|
||||
"output_cost_per_image_above_512_and_512_pixels_and_premium_image": {"type": "number"},
|
||||
"output_cost_per_image_above_512_and_512_pixels_and_premium_image": {
|
||||
"type": "number"
|
||||
},
|
||||
"output_cost_per_image_premium_image": {"type": "number"},
|
||||
"output_cost_per_token_batches": {"type": "number"},
|
||||
"output_cost_per_reasoning_token": {"type": "number"},
|
||||
|
|
@ -758,11 +852,15 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
|
|||
},
|
||||
}
|
||||
|
||||
prod_json = os.path.join(os.path.dirname(__file__), "..", "..", "model_prices_and_context_window.json")
|
||||
prod_json = os.path.join(
|
||||
os.path.dirname(__file__), "..", "..", "model_prices_and_context_window.json"
|
||||
)
|
||||
with open(prod_json, "r") as model_prices_file:
|
||||
actual_json = json.load(model_prices_file)
|
||||
assert isinstance(actual_json, dict)
|
||||
actual_json.pop("sample_spec", None) # remove the sample, whose schema is inconsistent with the real data
|
||||
actual_json.pop(
|
||||
"sample_spec", None
|
||||
) # remove the sample, whose schema is inconsistent with the real data
|
||||
|
||||
# Validate schema
|
||||
validate(actual_json, INTENDED_SCHEMA)
|
||||
|
|
@ -797,7 +895,9 @@ def test_max_tokens_consistency():
|
|||
from pathlib import Path
|
||||
|
||||
# Load the model configuration
|
||||
config_path = Path(__file__).parent.parent.parent / "model_prices_and_context_window.json"
|
||||
config_path = (
|
||||
Path(__file__).parent.parent.parent / "model_prices_and_context_window.json"
|
||||
)
|
||||
with open(config_path, "r") as f:
|
||||
models = json.load(f)
|
||||
|
||||
|
|
@ -827,9 +927,7 @@ def test_max_tokens_consistency():
|
|||
if inconsistencies:
|
||||
error_msg = f"\n\n❌ Found {len(inconsistencies)} models with max_tokens != max_output_tokens:\n\n"
|
||||
for item in inconsistencies[:10]: # Show first 10
|
||||
error_msg += (
|
||||
f" {item['model']}: max_tokens={item['max_tokens']}, max_output_tokens={item['max_output_tokens']}\n"
|
||||
)
|
||||
error_msg += f" {item['model']}: max_tokens={item['max_tokens']}, max_output_tokens={item['max_output_tokens']}\n"
|
||||
|
||||
if len(inconsistencies) > 10:
|
||||
error_msg += f"\n ... and {len(inconsistencies) - 10} more\n"
|
||||
|
|
@ -866,10 +964,15 @@ def test_openai_models_in_model_info():
|
|||
model_map = litellm.model_cost
|
||||
violated_models = []
|
||||
for model, info in model_map.items():
|
||||
if info.get("litellm_provider") == "openai" and info.get("supports_vision") is True:
|
||||
if (
|
||||
info.get("litellm_provider") == "openai"
|
||||
and info.get("supports_vision") is True
|
||||
):
|
||||
if info.get("supports_pdf_input") is not True:
|
||||
violated_models.append(model)
|
||||
assert len(violated_models) == 0, f"The following models should support pdf input: {violated_models}"
|
||||
assert (
|
||||
len(violated_models) == 0
|
||||
), f"The following models should support pdf input: {violated_models}"
|
||||
|
||||
|
||||
def test_supports_tool_choice_simple_tests():
|
||||
|
|
@ -877,8 +980,18 @@ def test_supports_tool_choice_simple_tests():
|
|||
simple sanity checks
|
||||
"""
|
||||
assert litellm.utils.supports_tool_choice(model="gpt-4o") == True
|
||||
assert litellm.utils.supports_tool_choice(model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0") == True
|
||||
assert litellm.utils.supports_tool_choice(model="anthropic.claude-3-sonnet-20240229-v1:0") is True
|
||||
assert (
|
||||
litellm.utils.supports_tool_choice(
|
||||
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
)
|
||||
== True
|
||||
)
|
||||
assert (
|
||||
litellm.utils.supports_tool_choice(
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0"
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
assert (
|
||||
litellm.utils.supports_tool_choice(
|
||||
|
|
@ -888,10 +1001,17 @@ def test_supports_tool_choice_simple_tests():
|
|||
is True
|
||||
)
|
||||
|
||||
assert litellm.utils.supports_tool_choice(model="us.amazon.nova-micro-v1:0") is False
|
||||
assert litellm.utils.supports_tool_choice(model="bedrock/us.amazon.nova-micro-v1:0") is False
|
||||
assert (
|
||||
litellm.utils.supports_tool_choice(model="us.amazon.nova-micro-v1:0", custom_llm_provider="bedrock_converse")
|
||||
litellm.utils.supports_tool_choice(model="us.amazon.nova-micro-v1:0") is False
|
||||
)
|
||||
assert (
|
||||
litellm.utils.supports_tool_choice(model="bedrock/us.amazon.nova-micro-v1:0")
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
litellm.utils.supports_tool_choice(
|
||||
model="us.amazon.nova-micro-v1:0", custom_llm_provider="bedrock_converse"
|
||||
)
|
||||
is False
|
||||
)
|
||||
|
||||
|
|
@ -984,7 +1104,9 @@ def test_supports_computer_use_utility():
|
|||
|
||||
try:
|
||||
# Test a model known to support computer_use from backup JSON
|
||||
supports_cu_anthropic = supports_computer_use(model="anthropic/claude-4-sonnet-20250514")
|
||||
supports_cu_anthropic = supports_computer_use(
|
||||
model="anthropic/claude-4-sonnet-20250514"
|
||||
)
|
||||
assert supports_cu_anthropic is True
|
||||
|
||||
# Test a model known not to have the flag or set to false (defaults to False via get_model_info)
|
||||
|
|
@ -1027,7 +1149,9 @@ def test_get_model_info_shows_supports_computer_use():
|
|||
model_known_not_to_support_computer_use = "gpt-3.5-turbo"
|
||||
info_gpt = litellm.get_model_info(model_known_not_to_support_computer_use)
|
||||
print(f"Info for {model_known_not_to_support_computer_use}: {info_gpt}")
|
||||
assert info_gpt.get("supports_computer_use") is None # Expecting None due to the default in ModelInfoBase
|
||||
assert (
|
||||
info_gpt.get("supports_computer_use") is None
|
||||
) # Expecting None due to the default in ModelInfoBase
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -1143,14 +1267,20 @@ class TestProxyFunctionCalling:
|
|||
("command-nightly", "litellm_proxy/command-nightly", False),
|
||||
],
|
||||
)
|
||||
def test_proxy_function_calling_support_consistency(self, direct_model, proxy_model, expected_result):
|
||||
def test_proxy_function_calling_support_consistency(
|
||||
self, direct_model, proxy_model, expected_result
|
||||
):
|
||||
"""Test that proxy models have the same function calling support as their direct counterparts."""
|
||||
direct_result = supports_function_calling(direct_model)
|
||||
proxy_result = supports_function_calling(proxy_model)
|
||||
|
||||
# Both should match the expected result
|
||||
assert direct_result == expected_result, f"Direct model {direct_model} should return {expected_result}"
|
||||
assert proxy_result == expected_result, f"Proxy model {proxy_model} should return {expected_result}"
|
||||
assert (
|
||||
direct_result == expected_result
|
||||
), f"Direct model {direct_model} should return {expected_result}"
|
||||
assert (
|
||||
proxy_result == expected_result
|
||||
), f"Proxy model {proxy_model} should return {expected_result}"
|
||||
|
||||
# Direct and proxy should be consistent
|
||||
assert (
|
||||
|
|
@ -1222,7 +1352,9 @@ class TestProxyFunctionCalling:
|
|||
("litellm_proxy/local-mistral", "ollama/mistral", False),
|
||||
],
|
||||
)
|
||||
def test_proxy_custom_model_names_without_config(self, proxy_model_name, underlying_model, expected_proxy_result):
|
||||
def test_proxy_custom_model_names_without_config(
|
||||
self, proxy_model_name, underlying_model, expected_proxy_result
|
||||
):
|
||||
"""
|
||||
Test proxy models with custom model names that differ from underlying models.
|
||||
|
||||
|
|
@ -1233,7 +1365,9 @@ class TestProxyFunctionCalling:
|
|||
# Test the underlying model directly first to establish what it SHOULD return
|
||||
try:
|
||||
underlying_result = supports_function_calling(underlying_model)
|
||||
print(f"Underlying model {underlying_model} supports function calling: {underlying_result}")
|
||||
print(
|
||||
f"Underlying model {underlying_model} supports function calling: {underlying_result}"
|
||||
)
|
||||
except Exception as e:
|
||||
print(f"Warning: Could not test underlying model {underlying_model}: {e}")
|
||||
|
||||
|
|
@ -1255,7 +1389,9 @@ class TestProxyFunctionCalling:
|
|||
# Case 1: Custom model name that cannot be resolved
|
||||
custom_model = "litellm_proxy/my-custom-claude"
|
||||
result = supports_function_calling(custom_model)
|
||||
assert result is False, "Custom model names return False without proxy config context"
|
||||
assert (
|
||||
result is False
|
||||
), "Custom model names return False without proxy config context"
|
||||
|
||||
# Case 2: Model name that can be resolved (matches pattern)
|
||||
resolvable_model = "litellm_proxy/claude-sonnet-4-5-20250929"
|
||||
|
|
@ -1294,7 +1430,9 @@ class TestProxyFunctionCalling:
|
|||
), # Hints at Bedrock Claude 3 Sonnet
|
||||
],
|
||||
)
|
||||
def test_proxy_models_with_naming_hints(self, proxy_model_with_hints, expected_result):
|
||||
def test_proxy_models_with_naming_hints(
|
||||
self, proxy_model_with_hints, expected_result
|
||||
):
|
||||
"""
|
||||
Test proxy models with names that provide hints about the underlying model.
|
||||
|
||||
|
|
@ -1306,10 +1444,14 @@ class TestProxyFunctionCalling:
|
|||
|
||||
# Currently these will return False, but we document the expected behavior
|
||||
# In the future, we could implement smarter model name inference
|
||||
print(f"Model {proxy_model_with_hints}: current={proxy_result}, desired={expected_result}")
|
||||
print(
|
||||
f"Model {proxy_model_with_hints}: current={proxy_result}, desired={expected_result}"
|
||||
)
|
||||
|
||||
# For now, we expect False (current behavior), but document the limitation
|
||||
assert proxy_result is False, f"Current limitation: {proxy_model_with_hints} returns False without inference"
|
||||
assert (
|
||||
proxy_result is False
|
||||
), f"Current limitation: {proxy_model_with_hints} returns False without inference"
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"proxy_model,expected_result",
|
||||
|
|
@ -1334,7 +1476,9 @@ class TestProxyFunctionCalling:
|
|||
"""
|
||||
try:
|
||||
result = supports_function_calling(model=proxy_model)
|
||||
assert result == expected_result, f"Proxy model {proxy_model} returned {result}, expected {expected_result}"
|
||||
assert (
|
||||
result == expected_result
|
||||
), f"Proxy model {proxy_model} returned {result}, expected {expected_result}"
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error testing proxy model {proxy_model}: {e}")
|
||||
|
||||
|
|
@ -1374,11 +1518,17 @@ class TestProxyFunctionCalling:
|
|||
parameter explicitly set to None, which is a common usage pattern.
|
||||
"""
|
||||
try:
|
||||
result = supports_function_calling(model=model_name, custom_llm_provider=None)
|
||||
result = supports_function_calling(
|
||||
model=model_name, custom_llm_provider=None
|
||||
)
|
||||
# All the models in this test should support function calling
|
||||
assert result is True, f"Model {model_name} should support function calling but returned {result}"
|
||||
assert (
|
||||
result is True
|
||||
), f"Model {model_name} should support function calling but returned {result}"
|
||||
except Exception as e:
|
||||
pytest.fail(f"Error testing {model_name} with custom_llm_provider=None: {e}")
|
||||
pytest.fail(
|
||||
f"Error testing {model_name} with custom_llm_provider=None: {e}"
|
||||
)
|
||||
|
||||
def test_edge_cases_and_malformed_proxy_models(self):
|
||||
"""Test edge cases and malformed proxy model names."""
|
||||
|
|
@ -1415,7 +1565,9 @@ class TestProxyFunctionCalling:
|
|||
proxy_result = supports_function_calling(model=proxy_model)
|
||||
|
||||
print("\nDemonstration of proxy model resolution:")
|
||||
print(f"Direct model '{direct_model}' supports function calling: {direct_result}")
|
||||
print(
|
||||
f"Direct model '{direct_model}' supports function calling: {direct_result}"
|
||||
)
|
||||
print(f"Proxy model '{proxy_model}' supports function calling: {proxy_result}")
|
||||
|
||||
# This assertion will currently fail due to the bug
|
||||
|
|
@ -1428,7 +1580,8 @@ class TestProxyFunctionCalling:
|
|||
)
|
||||
|
||||
assert direct_result == proxy_result, (
|
||||
f"Proxy model resolution issue: {direct_model} -> {direct_result}, " f"{proxy_model} -> {proxy_result}"
|
||||
f"Proxy model resolution issue: {direct_model} -> {direct_result}, "
|
||||
f"{proxy_model} -> {proxy_result}"
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -1605,7 +1758,9 @@ class TestProxyFunctionCalling:
|
|||
underlying_result is True
|
||||
), f"Claude 3 models should support function calling: {underlying_bedrock_model}"
|
||||
except Exception as e:
|
||||
print(f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}")
|
||||
print(
|
||||
f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}"
|
||||
)
|
||||
|
||||
# Test the proxy model - should return False due to lack of configuration context
|
||||
proxy_result = supports_function_calling(proxy_model_name)
|
||||
|
|
@ -1692,7 +1847,9 @@ class TestProxyFunctionCalling:
|
|||
result = supports_function_calling(model)
|
||||
print(f"Direct test - {model}: {result}")
|
||||
# Claude 3 models should support function calling
|
||||
assert result is True, f"Claude 3 model should support function calling: {model}"
|
||||
assert (
|
||||
result is True
|
||||
), f"Claude 3 model should support function calling: {model}"
|
||||
except Exception as e:
|
||||
print(f"Could not test {model}: {e}")
|
||||
|
||||
|
|
@ -1870,7 +2027,9 @@ class TestProxyFunctionCalling:
|
|||
underlying_result is True
|
||||
), f"Claude 3 models should support function calling: {underlying_bedrock_model}"
|
||||
except Exception as e:
|
||||
print(f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}")
|
||||
print(
|
||||
f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}"
|
||||
)
|
||||
|
||||
# Test the proxy model - should return False due to lack of configuration context
|
||||
proxy_result = supports_function_calling(proxy_model_name)
|
||||
|
|
@ -1957,7 +2116,9 @@ class TestProxyFunctionCalling:
|
|||
result = supports_function_calling(model)
|
||||
print(f"Direct test - {model}: {result}")
|
||||
# Claude 3 models should support function calling
|
||||
assert result is True, f"Claude 3 model should support function calling: {model}"
|
||||
assert (
|
||||
result is True
|
||||
), f"Claude 3 model should support function calling: {model}"
|
||||
except Exception as e:
|
||||
print(f"Could not test {model}: {e}")
|
||||
|
||||
|
|
@ -2135,7 +2296,9 @@ class TestProxyFunctionCalling:
|
|||
underlying_result is True
|
||||
), f"Claude 3 models should support function calling: {underlying_bedrock_model}"
|
||||
except Exception as e:
|
||||
print(f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}")
|
||||
print(
|
||||
f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}"
|
||||
)
|
||||
|
||||
# Test the proxy model - should return False due to lack of configuration context
|
||||
proxy_result = supports_function_calling(proxy_model_name)
|
||||
|
|
@ -2222,7 +2385,9 @@ class TestProxyFunctionCalling:
|
|||
result = supports_function_calling(model)
|
||||
print(f"Direct test - {model}: {result}")
|
||||
# Claude 3 models should support function calling
|
||||
assert result is True, f"Claude 3 model should support function calling: {model}"
|
||||
assert (
|
||||
result is True
|
||||
), f"Claude 3 model should support function calling: {model}"
|
||||
except Exception as e:
|
||||
print(f"Could not test {model}: {e}")
|
||||
|
||||
|
|
@ -2378,9 +2543,7 @@ def test_anthropic_claude_4_invoke_chat_provider_config():
|
|||
|
||||
def test_bedrock_application_inference_profile():
|
||||
model = "arn:aws:bedrock:us-east-2:<AWS-ACCOUNT-ID>:inference-profile/us.anthropic.claude-3-5-haiku-20241022-v1:0"
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm import completion
|
||||
from litellm.utils import supports_tool_choice
|
||||
|
||||
result = supports_tool_choice(model, custom_llm_provider="bedrock")
|
||||
|
|
@ -2447,7 +2610,6 @@ def test_block_key_hashing_logic():
|
|||
"""
|
||||
Test that block_key() function only hashes keys that start with "sk-"
|
||||
"""
|
||||
import hashlib
|
||||
|
||||
from litellm.proxy.utils import hash_token
|
||||
|
||||
|
|
@ -2473,13 +2635,17 @@ def test_block_key_hashing_logic():
|
|||
# Additional verification: if it should be hashed, verify it's actually a hash
|
||||
if should_be_hashed:
|
||||
# SHA-256 hashes are 64 characters long and contain only hex digits
|
||||
assert len(hashed_token) == 64, f"Hash length should be 64, got {len(hashed_token)} for {input_key}"
|
||||
assert (
|
||||
len(hashed_token) == 64
|
||||
), f"Hash length should be 64, got {len(hashed_token)} for {input_key}"
|
||||
assert all(
|
||||
c in "0123456789abcdef" for c in hashed_token
|
||||
), f"Hash should contain only hex digits for {input_key}"
|
||||
else:
|
||||
# If not hashed, it should be the original string
|
||||
assert hashed_token == input_key, f"Non-hashed key should remain unchanged: {input_key}"
|
||||
assert (
|
||||
hashed_token == input_key
|
||||
), f"Non-hashed key should remain unchanged: {input_key}"
|
||||
|
||||
print("✅ All block_key hashing logic tests passed!")
|
||||
|
||||
|
|
@ -2506,7 +2672,18 @@ def test_generate_gcp_iam_access_token():
|
|||
mock_iam_credentials_v1.GenerateAccessTokenRequest = Mock()
|
||||
|
||||
# Test successful token generation by mocking sys.modules
|
||||
with patch.dict("sys.modules", {"google.cloud.iam_credentials_v1": mock_iam_credentials_v1}):
|
||||
# Must mock parent packages too so `from google.cloud import iam_credentials_v1` works
|
||||
mock_google = Mock()
|
||||
mock_google_cloud = Mock()
|
||||
mock_google_cloud.iam_credentials_v1 = mock_iam_credentials_v1
|
||||
with patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"google": mock_google,
|
||||
"google.cloud": mock_google_cloud,
|
||||
"google.cloud.iam_credentials_v1": mock_iam_credentials_v1,
|
||||
},
|
||||
):
|
||||
from litellm._redis import _generate_gcp_iam_access_token
|
||||
|
||||
result = _generate_gcp_iam_access_token(service_account)
|
||||
|
|
@ -2557,13 +2734,22 @@ def test_generate_azure_ad_redis_token():
|
|||
mock_credential = Mock()
|
||||
mock_credential.get_token.return_value = mock_token
|
||||
|
||||
with patch("azure.identity.DefaultAzureCredential", return_value=mock_credential):
|
||||
mock_azure_identity = Mock()
|
||||
mock_azure_identity.DefaultAzureCredential = Mock(return_value=mock_credential)
|
||||
mock_azure_identity.ClientSecretCredential = Mock()
|
||||
mock_azure_identity.ManagedIdentityCredential = Mock()
|
||||
|
||||
with patch.dict(
|
||||
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
|
||||
):
|
||||
from litellm._redis import _generate_azure_ad_redis_token
|
||||
|
||||
result = _generate_azure_ad_redis_token()
|
||||
|
||||
assert result == expected_token
|
||||
mock_credential.get_token.assert_called_once_with("https://redis.azure.com/.default")
|
||||
mock_credential.get_token.assert_called_once_with(
|
||||
"https://redis.azure.com/.default"
|
||||
)
|
||||
|
||||
|
||||
def test_generate_azure_ad_redis_token_service_principal():
|
||||
|
|
@ -2578,7 +2764,16 @@ def test_generate_azure_ad_redis_token_service_principal():
|
|||
mock_credential = Mock()
|
||||
mock_credential.get_token.return_value = mock_token
|
||||
|
||||
with patch("azure.identity.ClientSecretCredential", return_value=mock_credential) as mock_cls:
|
||||
mock_client_secret_credential = Mock(return_value=mock_credential)
|
||||
|
||||
mock_azure_identity = Mock()
|
||||
mock_azure_identity.DefaultAzureCredential = Mock()
|
||||
mock_azure_identity.ClientSecretCredential = mock_client_secret_credential
|
||||
mock_azure_identity.ManagedIdentityCredential = Mock()
|
||||
|
||||
with patch.dict(
|
||||
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
|
||||
):
|
||||
from litellm._redis import _generate_azure_ad_redis_token
|
||||
|
||||
result = _generate_azure_ad_redis_token(
|
||||
|
|
@ -2588,7 +2783,7 @@ def test_generate_azure_ad_redis_token_service_principal():
|
|||
)
|
||||
|
||||
assert result == expected_token
|
||||
mock_cls.assert_called_once_with(
|
||||
mock_client_secret_credential.assert_called_once_with(
|
||||
client_id="test-client-id",
|
||||
tenant_id="test-tenant-id",
|
||||
client_secret="test-secret",
|
||||
|
|
@ -2616,31 +2811,23 @@ def test_generate_azure_ad_redis_token_import_error():
|
|||
|
||||
def test_redis_client_logic_azure_ad_auth():
|
||||
"""Test that _get_redis_client_logic sets up Azure AD auth when REDIS_AZURE_AD_TOKEN=true."""
|
||||
import os
|
||||
from unittest.mock import patch
|
||||
from litellm._redis import _get_redis_client_logic
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{
|
||||
"REDIS_HOST": "myredis.redis.cache.windows.net",
|
||||
"REDIS_PORT": "6380",
|
||||
"REDIS_AZURE_AD_TOKEN": "true",
|
||||
"REDIS_SSL": "true",
|
||||
},
|
||||
clear=True,
|
||||
):
|
||||
from litellm._redis import _get_redis_client_logic
|
||||
redis_kwargs = _get_redis_client_logic(
|
||||
host="myredis.redis.cache.windows.net",
|
||||
port="6380",
|
||||
azure_redis_ad_token="true",
|
||||
ssl=True,
|
||||
)
|
||||
|
||||
redis_kwargs = _get_redis_client_logic()
|
||||
# Should have redis_connect_func set
|
||||
assert "redis_connect_func" in redis_kwargs
|
||||
assert hasattr(redis_kwargs["redis_connect_func"], "_azure_redis_ad_token")
|
||||
assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True
|
||||
|
||||
# Should have redis_connect_func set
|
||||
assert "redis_connect_func" in redis_kwargs
|
||||
assert hasattr(redis_kwargs["redis_connect_func"], "_azure_redis_ad_token")
|
||||
assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True
|
||||
|
||||
# Azure-specific kwargs should be removed
|
||||
assert "azure_redis_ad_token" not in redis_kwargs
|
||||
assert "azure_client_id" not in redis_kwargs
|
||||
# Azure-specific kwargs should be removed
|
||||
assert "azure_redis_ad_token" not in redis_kwargs
|
||||
assert "azure_client_id" not in redis_kwargs
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
@ -2649,7 +2836,9 @@ if __name__ == "__main__":
|
|||
|
||||
|
||||
def test_model_info_for_vertex_ai_deepseek_model():
|
||||
model_info = litellm.get_model_info(model="vertex_ai/deepseek-ai/deepseek-r1-0528-maas")
|
||||
model_info = litellm.get_model_info(
|
||||
model="vertex_ai/deepseek-ai/deepseek-r1-0528-maas"
|
||||
)
|
||||
assert model_info is not None
|
||||
assert model_info["litellm_provider"] == "vertex_ai-deepseek_models"
|
||||
assert model_info["mode"] == "chat"
|
||||
|
|
@ -2679,7 +2868,9 @@ def test_model_info_for_openrouter_kimi_k2_5():
|
|||
model_cost = json.load(f)
|
||||
|
||||
model_info = model_cost.get("openrouter/moonshotai/kimi-k2.5")
|
||||
assert model_info is not None, "Model not found in model_prices_and_context_window.json"
|
||||
assert (
|
||||
model_info is not None
|
||||
), "Model not found in model_prices_and_context_window.json"
|
||||
assert model_info["litellm_provider"] == "openrouter"
|
||||
assert model_info["mode"] == "chat"
|
||||
|
||||
|
|
@ -2723,7 +2914,9 @@ def test_model_info_for_fireworks_short_form_models():
|
|||
"fireworks_ai/accounts/fireworks/models/glm-4p7",
|
||||
]:
|
||||
info = model_cost.get(key)
|
||||
assert info is not None, f"{key} not found in model_prices_and_context_window.json"
|
||||
assert (
|
||||
info is not None
|
||||
), f"{key} not found in model_prices_and_context_window.json"
|
||||
assert info["litellm_provider"] == "fireworks_ai"
|
||||
assert info["mode"] == "chat"
|
||||
assert info["input_cost_per_token"] == 6e-07
|
||||
|
|
@ -2737,7 +2930,9 @@ def test_model_info_for_fireworks_short_form_models():
|
|||
"fireworks_ai/accounts/fireworks/models/minimax-m2p1",
|
||||
]:
|
||||
info = model_cost.get(key)
|
||||
assert info is not None, f"{key} not found in model_prices_and_context_window.json"
|
||||
assert (
|
||||
info is not None
|
||||
), f"{key} not found in model_prices_and_context_window.json"
|
||||
assert info["litellm_provider"] == "fireworks_ai"
|
||||
assert info["mode"] == "chat"
|
||||
assert info["input_cost_per_token"] == 3e-07
|
||||
|
|
@ -2746,7 +2941,9 @@ def test_model_info_for_fireworks_short_form_models():
|
|||
|
||||
# kimi-k2p5: short-form only (long-form already existed)
|
||||
info = model_cost.get("fireworks_ai/kimi-k2p5")
|
||||
assert info is not None, "fireworks_ai/kimi-k2p5 not found in model_prices_and_context_window.json"
|
||||
assert (
|
||||
info is not None
|
||||
), "fireworks_ai/kimi-k2p5 not found in model_prices_and_context_window.json"
|
||||
assert info["litellm_provider"] == "fireworks_ai"
|
||||
assert info["mode"] == "chat"
|
||||
assert info["input_cost_per_token"] == 6e-07
|
||||
|
|
@ -2772,7 +2969,9 @@ class TestGetValidModelsWithCLI:
|
|||
]
|
||||
}
|
||||
|
||||
with patch.object(litellm.module_level_client, "get", return_value=mock_response) as mock_get:
|
||||
with patch.object(
|
||||
litellm.module_level_client, "get", return_value=mock_response
|
||||
) as mock_get:
|
||||
# Test the exact pattern used in cli_token_usage.py
|
||||
result = litellm.get_valid_models(
|
||||
check_provider_endpoint=True,
|
||||
|
|
@ -2958,7 +3157,9 @@ class TestProxyLoggingBudgetAlerts:
|
|||
user_info = MagicMock()
|
||||
|
||||
# Should not raise an error
|
||||
await proxy_logging.budget_alerts(type="organization_budget", user_info=user_info)
|
||||
await proxy_logging.budget_alerts(
|
||||
type="organization_budget", user_info=user_info
|
||||
)
|
||||
|
||||
async def test_budget_alerts_with_both_slack_and_email(self):
|
||||
"""Test that budget_alerts calls both slack and email instances when both are in alerting."""
|
||||
|
|
@ -3010,7 +3211,9 @@ class TestProxyLoggingBudgetAlerts:
|
|||
proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with(
|
||||
type=alert_type, user_info=user_info
|
||||
)
|
||||
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(type=alert_type, user_info=user_info)
|
||||
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
|
||||
type=alert_type, user_info=user_info
|
||||
)
|
||||
|
||||
async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_alerting_none(
|
||||
self,
|
||||
|
|
@ -3228,7 +3431,9 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926():
|
|||
{"role": "user", "content": "Build a feature"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"thinking_blocks": [{"type": "thinking", "thinking": "Let me analyze the requirements..."}],
|
||||
"thinking_blocks": [
|
||||
{"type": "thinking", "thinking": "Let me analyze the requirements..."}
|
||||
],
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "toolu_1",
|
||||
|
|
@ -3476,34 +3681,67 @@ class TestGetOptionalParamsDeepSeek:
|
|||
|
||||
class TestIsStreamingRequest:
|
||||
def test_stream_true_in_kwargs(self):
|
||||
assert _is_streaming_request(kwargs={"stream": True}, call_type="acompletion") is True
|
||||
assert (
|
||||
_is_streaming_request(kwargs={"stream": True}, call_type="acompletion")
|
||||
is True
|
||||
)
|
||||
|
||||
def test_stream_false_in_kwargs(self):
|
||||
assert _is_streaming_request(kwargs={"stream": False}, call_type="acompletion") is False
|
||||
assert (
|
||||
_is_streaming_request(kwargs={"stream": False}, call_type="acompletion")
|
||||
is False
|
||||
)
|
||||
|
||||
def test_no_stream_in_kwargs(self):
|
||||
assert _is_streaming_request(kwargs={}, call_type="acompletion") is False
|
||||
|
||||
def test_generate_content_stream_string(self):
|
||||
assert _is_streaming_request(kwargs={}, call_type=CallTypes.generate_content_stream.value) is True
|
||||
assert (
|
||||
_is_streaming_request(
|
||||
kwargs={}, call_type=CallTypes.generate_content_stream.value
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_agenerate_content_stream_string(self):
|
||||
assert _is_streaming_request(kwargs={}, call_type=CallTypes.agenerate_content_stream.value) is True
|
||||
assert (
|
||||
_is_streaming_request(
|
||||
kwargs={}, call_type=CallTypes.agenerate_content_stream.value
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_generate_content_stream_enum(self):
|
||||
assert _is_streaming_request(kwargs={}, call_type=CallTypes.generate_content_stream) is True
|
||||
assert (
|
||||
_is_streaming_request(
|
||||
kwargs={}, call_type=CallTypes.generate_content_stream
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_agenerate_content_stream_enum(self):
|
||||
assert _is_streaming_request(kwargs={}, call_type=CallTypes.agenerate_content_stream) is True
|
||||
assert (
|
||||
_is_streaming_request(
|
||||
kwargs={}, call_type=CallTypes.agenerate_content_stream
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_non_streaming_call_type_string(self):
|
||||
assert _is_streaming_request(kwargs={}, call_type="acompletion") is False
|
||||
|
||||
def test_non_streaming_call_type_enum(self):
|
||||
assert _is_streaming_request(kwargs={}, call_type=CallTypes.acompletion) is False
|
||||
assert (
|
||||
_is_streaming_request(kwargs={}, call_type=CallTypes.acompletion) is False
|
||||
)
|
||||
|
||||
def test_stream_true_overrides_non_streaming_call_type(self):
|
||||
assert _is_streaming_request(kwargs={"stream": True}, call_type=CallTypes.acompletion) is True
|
||||
assert (
|
||||
_is_streaming_request(
|
||||
kwargs={"stream": True}, call_type=CallTypes.acompletion
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
class TestCallbackAsyncSyncSeparation:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue