fix: req changes with relevant test cases

This commit is contained in:
Harshit28j 2026-03-07 05:49:03 +05:30
parent 3761598a02
commit 8bdbda0d2c
2 changed files with 444 additions and 149 deletions

View file

@ -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:

View file

@ -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: