From 8bdbda0d2c632ba0c939102024224737352919e2 Mon Sep 17 00:00:00 2001 From: Harshit28j Date: Sat, 7 Mar 2026 05:49:03 +0530 Subject: [PATCH] fix: req changes with relevant test cases --- litellm/_redis.py | 105 +++++-- tests/test_litellm/test_utils.py | 488 +++++++++++++++++++++++-------- 2 files changed, 444 insertions(+), 149 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index aa3689e8672..49098c0e17a 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -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: diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 70722c5095c..8bcceeb842c 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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::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: