diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py index 5d4e0d5fb39..404bcca712f 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py @@ -186,11 +186,11 @@ class BedrockAgentCoreA2ATransformation: headers: Final[dict] = {} session_id: Final = _validate_runtime_session_id( _request_scoped_runtime_session_id(params, litellm_params) - or agentcore_config._get_runtime_session_id(optional_params), + or agentcore_config.get_runtime_session_id(optional_params), model=model, ) headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = session_id - runtime_user_id: Final = agentcore_config._get_runtime_user_id(optional_params) + runtime_user_id: Final = agentcore_config.get_runtime_user_id(optional_params) if runtime_user_id: headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] = runtime_user_id diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 2b55e10a1d3..ba625ac5c00 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -369,7 +369,7 @@ def calculate_vertex_ai_batch_cost_and_usage( row, model_name, model_info=model_info, - calculate_usage=VertexGeminiConfig._calculate_usage, + calculate_usage=VertexGeminiConfig.calculate_usage, cost_calculator=batch_cost_calculator, ) for row in vertex_ai_batch_responses diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 3ba27cc9791..548cf693ab9 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -628,7 +628,7 @@ def retrieve_batch( async_kwargs: Final = kwargs.copy() async_kwargs.pop("aws_region_name", None) - return BedrockBatchesHandler._handle_async_invoke_status( + return BedrockBatchesHandler.handle_async_invoke_status( batch_id=batch_id, aws_region_name=kwargs.get("aws_region_name", "us-east-1"), logging_obj=litellm_logging_obj, @@ -638,7 +638,7 @@ def retrieve_batch( mij_kwargs: Final = kwargs.copy() mij_kwargs.pop("aws_region_name", None) - return BedrockBatchesHandler._handle_model_invocation_job_status( + return BedrockBatchesHandler.handle_model_invocation_job_status( batch_id=batch_id, aws_region_name=kwargs.get("aws_region_name"), logging_obj=litellm_logging_obj, @@ -1107,7 +1107,7 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj embedding_handler: Final = BedrockEmbedding() # Get the status of the async invoke job - status_response: Final = await embedding_handler._get_async_invoke_status( + status_response: Final = await embedding_handler.get_async_invoke_status( invocation_arn=batch_id, aws_region_name=aws_region_name, logging_obj=logging_obj, @@ -1151,7 +1151,7 @@ def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj failed_at, _, _, - ) = BedrockBatchesConfig()._parse_timestamps_and_status(status_response, aws_status_raw) + ) = BedrockBatchesConfig().parse_timestamps_and_status(status_response, aws_status_raw) result: Final = LiteLLMBatch( id=status_response["invocationArn"], object="batch", diff --git a/litellm/caching/gcs_cache.py b/litellm/caching/gcs_cache.py index e9922218828..76fcba86eb0 100644 --- a/litellm/caching/gcs_cache.py +++ b/litellm/caching/gcs_cache.py @@ -10,8 +10,8 @@ from urllib.parse import quote from litellm._logging import print_verbose, verbose_logger from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) @@ -31,7 +31,7 @@ class GCSCache(BaseCache): self.key_prefix = gcs_path.rstrip("/") + "/" if gcs_path else "" # create httpx clients self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) - self.sync_client = _get_httpx_client() + self.sync_client = get_httpx_client() def _construct_headers(self) -> dict: base: Final = GCSBucketBase(bucket_name=self.bucket_name) diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index b99023c07fd..8dac3f2eef9 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -68,8 +68,8 @@ class QdrantSemanticCache(BaseCache): embedding_timeout: float | None = None, ): from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) from litellm.secret_managers.main import get_secret_str @@ -112,7 +112,7 @@ class QdrantSemanticCache(BaseCache): self.headers = headers - self.sync_client = _get_httpx_client() + self.sync_client = get_httpx_client() self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Caching) if quantization_config is None: diff --git a/litellm/compression/compress.py b/litellm/compression/compress.py index bc2cb7c9fdc..6a3e44bda65 100644 --- a/litellm/compression/compress.py +++ b/litellm/compression/compress.py @@ -62,7 +62,7 @@ def _build_retrieval_tools(keys: list[str], call_type: str) -> list[dict]: # module import for non-Anthropic call paths. from litellm.llms.anthropic.chat.transformation import AnthropicConfig - anthropic_tools, _mcp_servers = AnthropicConfig()._map_tools(openai_tools) + anthropic_tools, _mcp_servers = AnthropicConfig().map_tools(openai_tools) return cast(list[dict], anthropic_tools) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index b637676836c..4d5f882ea10 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -70,15 +70,15 @@ from litellm.llms.gemini.cost_calculator import cost_per_token as gemini_cost_pe from litellm.llms.lemonade.cost_calculator import ( cost_per_token as lemonade_cost_per_token, ) -from litellm.llms.openai.cost_calculation import ( - _video_output_cost_per_second, -) from litellm.llms.openai.cost_calculation import ( cost_per_second as openai_cost_per_second, ) from litellm.llms.openai.cost_calculation import ( cost_per_token as openai_cost_per_token, ) +from litellm.llms.openai.cost_calculation import ( + video_output_cost_per_second, +) from litellm.llms.perplexity.cost_calculator import ( cost_per_token as perplexity_cost_per_token, ) @@ -2552,7 +2552,7 @@ def default_video_cost_calculator( if video_cost_per_second is not None: return video_cost_per_second * duration_seconds - output_cost_per_second: Final = _video_output_cost_per_second(cost_info, video_resolution) + output_cost_per_second: Final = video_output_cost_per_second(cost_info, video_resolution) if output_cost_per_second is not None: return output_cost_per_second * duration_seconds diff --git a/litellm/decisions/main.py b/litellm/decisions/main.py index 44a685e3235..8304c52a886 100644 --- a/litellm/decisions/main.py +++ b/litellm/decisions/main.py @@ -10,7 +10,7 @@ import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.decisions.transformation import DecisionsProviderConfig from litellm.llms.cloudflare.decisions.transformation import CLOUDFLARE_DECISIONS_ENDPOINT -from litellm.llms.custom_httpx.http_handler import _get_httpx_client, get_async_httpx_client +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, get_httpx_client from litellm.llms.openrouter.decisions.transformation import OPENROUTER_DECISIONS_ENDPOINT from litellm.llms.perplexity.decisions.transformation import PERPLEXITY_DECISIONS_ENDPOINT from litellm.llms.strands_decider.decisions.transformation import STRANDS_DECIDER_DECISIONS_ENDPOINT @@ -283,7 +283,7 @@ def decisions( ) logging_obj: Final = _log_request(prepared, kwargs) try: - handler: Final = _get_httpx_client() + handler: Final = get_httpx_client() response: Final = handler.post( prepared.url, json=dict(prepared.body), diff --git a/litellm/integrations/agentops/agentops.py b/litellm/integrations/agentops/agentops.py index 399ee49238c..9043ab1853a 100644 --- a/litellm/integrations/agentops/agentops.py +++ b/litellm/integrations/agentops/agentops.py @@ -7,7 +7,7 @@ from dataclasses import dataclass from typing import Any, Final from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client @dataclass @@ -99,7 +99,7 @@ class AgentOps(OpenTelemetry): "Connection": "keep-alive", } - client: Final = _get_httpx_client() + client: Final = get_httpx_client() try: response: Final = client.post( url=auth_endpoint, diff --git a/litellm/integrations/datadog/datadog.py b/litellm/integrations/datadog/datadog.py index d70acd51679..7046b8d017a 100644 --- a/litellm/integrations/datadog/datadog.py +++ b/litellm/integrations/datadog/datadog.py @@ -44,8 +44,8 @@ from litellm.integrations.datadog.datadog_mock_client import ( ) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus @@ -180,7 +180,7 @@ class DataDogLogger( dd_base_url: Final = get_datadog_base_url_from_env() if dd_base_url: self.intake_url = f"{dd_base_url}/api/v2/logs" - self.sync_client = _get_httpx_client() + self.sync_client = get_httpx_client() asyncio.create_task(self.periodic_flush()) self.flush_lock = asyncio.Lock() super().__init__( diff --git a/litellm/integrations/gcs_bucket/gcs_bucket_base.py b/litellm/integrations/gcs_bucket/gcs_bucket_base.py index 2f7fc4a52d9..0856a420d6f 100644 --- a/litellm/integrations/gcs_bucket/gcs_bucket_base.py +++ b/litellm/integrations/gcs_bucket/gcs_bucket_base.py @@ -53,13 +53,13 @@ class GCSBucketBase(CustomBatchLogger): if vertex_instance is None: vertex_instance = vertex_chat_completion - _auth_header, vertex_project = await vertex_instance._ensure_access_token_async( + _auth_header, vertex_project = await vertex_instance.ensure_access_token_async( credentials=service_account_json, project_id=None, custom_llm_provider="vertex_ai", ) - auth_header, _ = vertex_instance._get_token_and_url( + auth_header, _ = vertex_instance.get_token_and_url( model="gcs-bucket", auth_header=_auth_header, vertex_credentials=service_account_json, @@ -89,13 +89,13 @@ class GCSBucketBase(CustomBatchLogger): # from Secret Manager. project_id: Final = os.getenv("GOOGLE_SECRET_MANAGER_PROJECT_ID") - _auth_header, vertex_project = vertex_chat_completion._ensure_access_token( + _auth_header, vertex_project = vertex_chat_completion.ensure_access_token( credentials=self.path_service_account_json, project_id=project_id, custom_llm_provider="vertex_ai", ) - auth_header, _ = vertex_chat_completion._get_token_and_url( + auth_header, _ = vertex_chat_completion.get_token_and_url( model="gcs-bucket", auth_header=_auth_header, vertex_credentials=self.path_service_account_json, @@ -192,7 +192,7 @@ class GCSBucketBase(CustomBatchLogger): _in_memory_key: Final = self._get_in_memory_key_for_vertex_instance(credentials) if _in_memory_key not in self.vertex_instances: vertex_instance: Final = VertexBase() - await vertex_instance._ensure_access_token_async( + await vertex_instance.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", diff --git a/litellm/integrations/gcs_pubsub/pub_sub.py b/litellm/integrations/gcs_pubsub/pub_sub.py index 3e14d903e63..293be174811 100644 --- a/litellm/integrations/gcs_pubsub/pub_sub.py +++ b/litellm/integrations/gcs_pubsub/pub_sub.py @@ -69,13 +69,13 @@ class GcsPubSubLogger(CustomBatchLogger): ( _auth_header, vertex_project, - ) = await vertex_chat_completion._ensure_access_token_async( + ) = await vertex_chat_completion.ensure_access_token_async( credentials=self.path_service_account_json, project_id=self.project_id, custom_llm_provider="vertex_ai", ) - auth_header, _ = vertex_chat_completion._get_token_and_url( + auth_header, _ = vertex_chat_completion.get_token_and_url( model="pub-sub", auth_header=_auth_header, vertex_credentials=self.path_service_account_json, diff --git a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py index d78d1e7d223..2601f44e0bf 100644 --- a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py +++ b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py @@ -15,8 +15,8 @@ from litellm.integrations.prompt_management_base import ( PromptManagementClient, ) from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.llms.openai import AllMessageValues @@ -114,7 +114,7 @@ class GenericPromptManager(CustomPromptManagement): "prompt_id": prompt_id, **(self.additional_provider_specific_query_params or {}), } - http_client: Final = _get_httpx_client() + http_client: Final = get_httpx_client() try: response: Final = http_client.get( diff --git a/litellm/integrations/humanloop.py b/litellm/integrations/humanloop.py index 9e52ccd3c02..093a68c9284 100644 --- a/litellm/integrations/humanloop.py +++ b/litellm/integrations/humanloop.py @@ -11,7 +11,7 @@ from typing_extensions import TypedDict import litellm from litellm.caching import DualCache -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues from litellm.types.prompts.init_prompts import PromptSpec @@ -61,7 +61,7 @@ class HumanLoopPromptManager(DualCache): return compiled_prompts def _get_prompt_from_id_api(self, humanloop_prompt_id: str, humanloop_api_key: str) -> PromptManagementClient: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() base_url: Final = f"https://api.humanloop.com/v5/prompts/{humanloop_prompt_id}" diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 126ea1da728..faba4f5eb20 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -28,7 +28,7 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( validate_langfuse_environment_value, ) from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client from litellm.secret_managers.main import str_to_bool from litellm.types.integrations.langfuse import * from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse @@ -324,7 +324,7 @@ class LangFuseLogger: self.langfuse_client = create_mock_langfuse_client() self.is_mock_mode = True else: - self._http_handler: Final = _get_httpx_client() + self._http_handler: Final = get_httpx_client() self.langfuse_client = self._http_handler.client self.is_mock_mode = False diff --git a/litellm/integrations/langfuse/langfuse_sdk.py b/litellm/integrations/langfuse/langfuse_sdk.py index d117120dd70..864c0697eb3 100644 --- a/litellm/integrations/langfuse/langfuse_sdk.py +++ b/litellm/integrations/langfuse/langfuse_sdk.py @@ -40,7 +40,7 @@ import litellm from litellm._logging import verbose_logger from litellm.integrations.langfuse.langfuse import PROMPT_CACHE_TTL_ENV, parse_langfuse_debug, whole_number from litellm.litellm_core_utils.safe_json_dumps import safe_dumps -from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client +from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_httpx_client from litellm.types.llms.base import LiteLLMBaseModel __all__ = ( @@ -793,7 +793,7 @@ def _build_span_exporter(*, public_key: str, secret_key: str, base_url: str) -> export_path: Final = os.getenv("LANGFUSE_OTEL_TRACES_EXPORT_PATH") or "/api/public/otel/v1/traces" encoded_auth: Final = b64encode(f"{public_key}:{secret_key}".encode()).decode("ascii") return LangfuseSpanExporter( - handler=_get_httpx_client(), + handler=get_httpx_client(), endpoint=f"{base_url.rstrip('/')}/{export_path.lstrip('/')}", headers=MappingProxyType( { diff --git a/litellm/integrations/opik/opik.py b/litellm/integrations/opik/opik.py index ce47d7fe27a..80d8384310a 100644 --- a/litellm/integrations/opik/opik.py +++ b/litellm/integrations/opik/opik.py @@ -13,8 +13,8 @@ from typing_extensions import ReadOnly, TypedDict, Unpack from litellm._logging import verbose_logger from litellm.integrations.custom_batch_logger import CustomBatchLogger from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) @@ -53,7 +53,7 @@ class OpikLogger(CustomBatchLogger): def __init__(self, **kwargs: Unpack[_OpikLoggerKwargs]) -> None: self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) - self.sync_httpx_client = _get_httpx_client() + self.sync_httpx_client = get_httpx_client() self.opik_project_name: str = ( utils.get_opik_config_variable( diff --git a/litellm/integrations/posthog.py b/litellm/integrations/posthog.py index 4f7dff952e6..0f165218321 100644 --- a/litellm/integrations/posthog.py +++ b/litellm/integrations/posthog.py @@ -26,8 +26,8 @@ from litellm.integrations.posthog_mock_client import ( ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) from litellm.types.integrations.posthog import ( @@ -73,7 +73,7 @@ class PostHogLogger(CustomBatchLogger): raise Exception("POSTHOG_API_KEY is not set, set 'POSTHOG_API_KEY=<>'") self.async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) - self.sync_client = _get_httpx_client() + self.sync_client = get_httpx_client() self.POSTHOG_API_KEY = os.getenv("POSTHOG_API_KEY") posthog_api_url: Final = os.getenv("POSTHOG_API_URL", "https://us.i.posthog.com") diff --git a/litellm/integrations/s3_v2.py b/litellm/integrations/s3_v2.py index 08249d33561..47d9b1c7b64 100644 --- a/litellm/integrations/s3_v2.py +++ b/litellm/integrations/s3_v2.py @@ -52,8 +52,8 @@ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) from litellm.types.integrations.s3_v2 import S3PartitionGranularity, s3BatchLoggingElement @@ -875,7 +875,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM): prepared: Final = self._prepare_put(batch_logging_element) - httpx_client: Final = _get_httpx_client( + httpx_client: Final = get_httpx_client( params=({"ssl_verify": self.s3_verify} if self.s3_verify is not None else None) ) diff --git a/litellm/interactions/http_handler.py b/litellm/interactions/http_handler.py index 17ec4a3398d..f02e7e7679d 100644 --- a/litellm/interactions/http_handler.py +++ b/litellm/interactions/http_handler.py @@ -20,8 +20,8 @@ from litellm.llms.base_llm.interactions.transformation import BaseInteractionsAP from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.interactions import ( CancelInteractionResult, @@ -57,7 +57,7 @@ class _BaseHTTPHandler: litellm_params: GenericLiteLLMParams, client: HTTPHandler | None, ) -> HTTPHandler: - return client or _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + return client or get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) def _async_client( self, @@ -129,7 +129,7 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): ) if client is None: - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -363,7 +363,7 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): ) if client is None: - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -482,7 +482,7 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): ) if client is None: - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -603,7 +603,7 @@ class InteractionsHTTPHandler(_BaseHTTPHandler): ) if client is None: - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index 5a498ead6bc..724be5510f5 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -119,7 +119,7 @@ def handle_anthropic_text_model_custom_llm_provider( """ if custom_llm_provider: - if custom_llm_provider == "anthropic" and litellm.AnthropicTextConfig._is_anthropic_text_model(model): + if custom_llm_provider == "anthropic" and litellm.AnthropicTextConfig.is_anthropic_text_model(model): return model, "anthropic_text" if model and "/" in model: @@ -127,7 +127,7 @@ def handle_anthropic_text_model_custom_llm_provider( if ( _custom_llm_provider and _custom_llm_provider == "anthropic" - and litellm.AnthropicTextConfig._is_anthropic_text_model(_model) + and litellm.AnthropicTextConfig.is_anthropic_text_model(_model) ): return _model, "anthropic_text" @@ -426,7 +426,7 @@ def get_llm_provider( custom_llm_provider = "text-completion-openai" ## anthropic elif model in litellm.anthropic_models: - if litellm.AnthropicTextConfig._is_anthropic_text_model(model): + if litellm.AnthropicTextConfig.is_anthropic_text_model(model): custom_llm_provider = "anthropic_text" else: custom_llm_provider = "anthropic" @@ -605,7 +605,7 @@ def _get_openai_compatible_provider_info( if provider_config is None: raise ValueError(f"Provider {custom_llm_provider} not found") config_class: Final = create_config_class(provider_config) - api_base, dynamic_api_key = config_class()._get_openai_compatible_provider_info(api_base, api_key) + api_base, dynamic_api_key = config_class().get_openai_compatible_provider_info(api_base, api_key) return model, custom_llm_provider, dynamic_api_key, api_base if custom_llm_provider == "perplexity": @@ -613,7 +613,7 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.PerplexityChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.PerplexityChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "aiohttp_openai": return model, "aiohttp_openai", api_key, api_base elif custom_llm_provider == "anyscale": @@ -624,7 +624,7 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.DeepInfraConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.DeepInfraConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "empower": api_base = api_base or get_secret("EMPOWER_API_BASE") or "https://app.empower.dev/api/v1" dynamic_api_key = api_key or get_secret_str("EMPOWER_API_KEY") @@ -632,14 +632,14 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.GroqChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.GroqChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "bedrock_mantle": from litellm.llms.bedrock_mantle.common_utils import split_mantle_region_prefix ( api_base, dynamic_api_key, - ) = litellm.BedrockMantleChatConfig()._get_openai_compatible_provider_info( + ) = litellm.BedrockMantleChatConfig().get_openai_compatible_provider_info( api_base, api_key, litellm_params=litellm_params, model=model ) model = split_mantle_region_prefix(model)[1] # rebind-ok: the prefix is routing only, not a Mantle model id @@ -704,25 +704,25 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.HostedVLLMChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.HostedVLLMChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "llamafile": # llamafile is OpenAI compatible. ( api_base, dynamic_api_key, - ) = litellm.LlamafileChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.LlamafileChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "datarobot": # DataRobot is OpenAI compatible. ( api_base, dynamic_api_key, - ) = litellm.DataRobotConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.DataRobotConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "lm_studio": # lm_studio is openai compatible, we just need to set this to custom_openai ( api_base, dynamic_api_key, - ) = litellm.LMStudioChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.LMStudioChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "deepseek": # deepseek is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.deepseek.com/v1 api_base = api_base or get_secret("DEEPSEEK_API_BASE") or "https://api.deepseek.com/beta" @@ -737,13 +737,13 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.FireworksAIConfig()._get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) + ) = litellm.FireworksAIConfig().get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) elif custom_llm_provider == "azure_ai": ( api_base, dynamic_api_key, custom_llm_provider, - ) = litellm.AzureAIStudioConfig()._get_openai_compatible_provider_info( + ) = litellm.AzureAIStudioConfig().get_openai_compatible_provider_info( model, api_base, api_key, custom_llm_provider ) elif custom_llm_provider == "github": @@ -757,29 +757,29 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.LiteLLMProxyChatConfig()._get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) + ) = litellm.LiteLLMProxyChatConfig().get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) elif custom_llm_provider == "mistral": ( api_base, dynamic_api_key, - ) = litellm.MistralConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.MistralConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "jina_ai": ( custom_llm_provider, api_base, dynamic_api_key, - ) = litellm.JinaAIEmbeddingConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.JinaAIEmbeddingConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "xai": ( api_base, dynamic_api_key, - ) = litellm.XAIChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.XAIChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "zai": ( api_base, dynamic_api_key, - ) = litellm.ZAIChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.ZAIChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "together_ai": api_base = api_base or get_secret_str("TOGETHER_AI_API_BASE") or "https://api.together.ai/v1" dynamic_api_key = api_key or ( @@ -799,7 +799,7 @@ def _get_openai_compatible_provider_info( api_base, dynamic_api_key, custom_llm_provider, - ) = litellm.GithubCopilotConfig()._get_openai_compatible_provider_info( + ) = litellm.GithubCopilotConfig().get_openai_compatible_provider_info( model, api_base, api_key, custom_llm_provider ) elif custom_llm_provider == "chatgpt": @@ -807,7 +807,7 @@ def _get_openai_compatible_provider_info( api_base, dynamic_api_key, custom_llm_provider, - ) = litellm.ChatGPTConfig()._get_openai_compatible_provider_info(model, api_base, api_key, custom_llm_provider) + ) = litellm.ChatGPTConfig().get_openai_compatible_provider_info(model, api_base, api_key, custom_llm_provider) elif custom_llm_provider == "novita": api_base = api_base or get_secret("NOVITA_API_BASE") or "https://api.novita.ai/v3/openai" dynamic_api_key = api_key or get_secret_str("NOVITA_API_KEY") @@ -815,78 +815,78 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.SnowflakeConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.SnowflakeConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "gradient_ai": ( api_base, dynamic_api_key, - ) = litellm.GradientAIConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.GradientAIConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "featherless_ai": ( api_base, dynamic_api_key, - ) = litellm.FeatherlessAIConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.FeatherlessAIConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "nscale": ( api_base, dynamic_api_key, - ) = litellm.NscaleConfig()._get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) + ) = litellm.NscaleConfig().get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) elif custom_llm_provider == "heroku": ( api_base, dynamic_api_key, - ) = litellm.HerokuChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.HerokuChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider in ("dashscope", "qwencloud", "qwen_ai_platform"): ( api_base, dynamic_api_key, - ) = _dashscope_family_chat_config(custom_llm_provider)._get_openai_compatible_provider_info(api_base, api_key) + ) = _dashscope_family_chat_config(custom_llm_provider).get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "modelscope": ( api_base, dynamic_api_key, - ) = litellm.ModelScopeChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.ModelScopeChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "moonshot": ( api_base, dynamic_api_key, - ) = litellm.MoonshotChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.MoonshotChatConfig().get_openai_compatible_provider_info(api_base, api_key) # publicai is now handled by JSON config (see litellm/llms/openai_like/providers.json) elif custom_llm_provider == "docker_model_runner": ( api_base, dynamic_api_key, - ) = litellm.DockerModelRunnerChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.DockerModelRunnerChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "v0": ( api_base, dynamic_api_key, - ) = litellm.V0ChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.V0ChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "morph": ( api_base, dynamic_api_key, - ) = litellm.MorphChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.MorphChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "lambda_ai": ( api_base, dynamic_api_key, - ) = litellm.LambdaAIChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.LambdaAIChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "inception": ( api_base, dynamic_api_key, - ) = litellm.InceptionChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.InceptionChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "hyperbolic": ( api_base, dynamic_api_key, - ) = litellm.HyperbolicChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.HyperbolicChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "vercel_ai_gateway": ( api_base, dynamic_api_key, - ) = litellm.VercelAIGatewayConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.VercelAIGatewayConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "edenai": api_base = litellm.EdenAIChatConfig.get_api_base(api_base) # rebind-ok: chain resolves in place dynamic_api_key = litellm.EdenAIChatConfig.get_api_key(api_key) # rebind-ok: chain resolves in place @@ -896,7 +896,7 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.AIMLChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.AIMLChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "wandb": api_base = api_base or get_secret("WANDB_API_BASE") or "https://api.inference.wandb.ai/v1" dynamic_api_key = api_key or get_secret_str("WANDB_API_KEY") @@ -904,19 +904,19 @@ def _get_openai_compatible_provider_info( ( api_base, dynamic_api_key, - ) = litellm.LemonadeChatConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.LemonadeChatConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "clarifai": ( api_base, dynamic_api_key, - ) = litellm.ClarifaiConfig()._get_openai_compatible_provider_info(api_base, api_key) + ) = litellm.ClarifaiConfig().get_openai_compatible_provider_info(api_base, api_key) elif custom_llm_provider == "ragflow": full_model: Final = f"ragflow/{model}" ( api_base, dynamic_api_key, _, - ) = litellm.RAGFlowConfig()._get_openai_compatible_provider_info(full_model, api_base, api_key, "ragflow") + ) = litellm.RAGFlowConfig().get_openai_compatible_provider_info(full_model, api_base, api_key, "ragflow") model = full_model elif custom_llm_provider == "langgraph": # LangGraph is a custom provider, just need to set api_base diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e15447c8584..2ce05b9aa39 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -2116,7 +2116,7 @@ class Logging(LiteLLMLoggingBaseClass): import httpx completion_response = result.model_dump(by_alias=True) if isinstance(result, BaseModel) else dict(result) - return litellm.VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response( + return litellm.VertexGeminiConfig().transform_google_generate_content_to_openai_model_response( completion_response=completion_response, model_response=ModelResponse(), model=self.model or "", @@ -4361,7 +4361,7 @@ class Logging(LiteLLMLoggingBaseClass): if httpx_response is None: raise ValueError("Google GenAI Generate Content: httpx_response is None") dict_result: Final = httpx_response.json() - result = litellm.VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response( + result = litellm.VertexGeminiConfig().transform_google_generate_content_to_openai_model_response( completion_response=dict_result, model_response=litellm.ModelResponse(), model=self.model, diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 017a75dca2e..973a2787b6a 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1303,7 +1303,7 @@ def convert_to_gemini_tool_call_invoke( VertexGeminiConfig, ) - needs_dummy_signature: Final = model is not None and VertexGeminiConfig._is_gemini_3_or_newer(model) + needs_dummy_signature: Final = model is not None and VertexGeminiConfig.is_gemini_3_or_newer(model) if tool_calls is not None: for tool in tool_calls: @@ -5312,7 +5312,7 @@ def prompt_factory( if custom_llm_provider == "ollama": return ollama_pt(model=model, messages=messages) elif custom_llm_provider == "anthropic": - if litellm.AnthropicTextConfig._is_anthropic_text_model(model): + if litellm.AnthropicTextConfig.is_anthropic_text_model(model): return anthropic_pt(messages=messages) return anthropic_messages_pt(messages=messages, model=model, llm_provider=custom_llm_provider) elif custom_llm_provider == "anthropic_xml": @@ -5327,7 +5327,7 @@ def prompt_factory( else: return gemini_text_image_pt(messages=messages) elif custom_llm_provider == "mistral": - return litellm.MistralConfig()._transform_messages(messages=messages, model=model) + return litellm.MistralConfig().transform_messages(messages=messages, model=model) elif custom_llm_provider == "bedrock": if "amazon.titan-text" in model: return amazon_titan_pt(messages=messages) diff --git a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py index 1f54f62392d..6ab65d5c09c 100644 --- a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py +++ b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py @@ -5,8 +5,8 @@ from typing import Any, Final, Literal from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.custom_http import httpxSpecialProvider @@ -50,7 +50,7 @@ def get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: """ try: url: Final = f"https://huggingface.co/{hf_model_name}/raw/main/tokenizer_config.json" - client: Final = _get_httpx_client() + client: Final = get_httpx_client() response: Final = client.get(url=url) except Exception as e: raise e @@ -103,7 +103,7 @@ def get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: Dict with 'status' and optionally 'chat_template' keys """ template_filenames: Final = ["chat_template.jinja", "chat_template.jinja2"] - client: Final = _get_httpx_client() + client: Final = get_httpx_client() for filename in template_filenames: try: diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 1027a681d56..43611eb3192 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1315,7 +1315,7 @@ class CustomStreamWrapper: raise ValueError(f"chunk is not a string: {chunk}") response_obj = cast( dict[str, object], - litellm.CodestralTextCompletionConfig()._chunk_parser(chunk), + litellm.CodestralTextCompletionConfig().chunk_parser(chunk), ) completion_obj["content"] = response_obj["text"] print_verbose(f"completion obj content: {completion_obj['content']}") diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 7357eb5ef19..4286d4d6f6a 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -32,7 +32,7 @@ from litellm.constants import ( from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.tokenizer import Encoding, HuggingFace, HuggingFaceTokenizer, OpenAIEncoding from litellm.litellm_core_utils.url_utils import safe_get -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client from litellm.rust_bridge.tokenizer import get_encoding from litellm.types.llms.anthropic import ( AnthropicContentParamSource, @@ -232,7 +232,7 @@ def get_image_dimensions( img_data = None if data.startswith(("http://", "https://")): try: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() response: Final[httpx.Response] = safe_get(client, data) max_bytes: Final = int(MAX_IMAGE_URL_DOWNLOAD_SIZE_MB * 1024 * 1024) content_length: Final[str | None] = response.headers.get("Content-Length") diff --git a/litellm/llms/aiml/chat/transformation.py b/litellm/llms/aiml/chat/transformation.py index 55bd754fd40..c9f8786f3be 100644 --- a/litellm/llms/aiml/chat/transformation.py +++ b/litellm/llms/aiml/chat/transformation.py @@ -18,3 +18,10 @@ class AIMLChatConfig(OpenAIGPTConfig): ) dynamic_api_key: Final = api_key or get_secret_str("AIML_API_KEY") return api_base, dynamic_api_key + + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 578a4e47dfd..405ebce548a 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -633,7 +633,7 @@ class AnthropicMessagesHandler(BaseTranslation): anthropic_config: Final = AnthropicConfig() anthropic_tools: Final[list[AllAnthropicToolsValues]] = [] for tool in guardrailed_tools: - converted_tool, mcp_server = anthropic_config._map_tool_helper(tool) + converted_tool, mcp_server = anthropic_config.map_tool_helper(tool) if converted_tool is not None: anthropic_tools.append(converted_tool) # Note: MCP servers are handled separately in the main transformation diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 24c8bc07c76..a110cc5b15d 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -23,8 +23,8 @@ from litellm.litellm_core_utils.json_fragment_accumulator import JSONFragmentAcc from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.anthropic import ( ContentBlockDelta, @@ -496,7 +496,7 @@ class AnthropicChatCompletion(BaseLLM): else: if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={"timeout": timeout}) + client = get_httpx_client(params={"timeout": timeout}) else: client = client @@ -1074,7 +1074,7 @@ class ModelResponseIterator: # Convert tool to content if we're tracking a response_format tool if self.is_response_format_tool: - message: Final = AnthropicConfig._convert_tool_response_to_message(tool_calls=[tool_use]) + message: Final = AnthropicConfig.convert_tool_response_to_message(tool_calls=[tool_use]) if message is not None: text = message.content or "" tool_use = None diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 96e181b7dcf..54a8780ca5b 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -416,6 +416,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return f"effort='xhigh' is not supported by this model. Got model: {model}" return None + @classmethod + def validate_effort_for_model( + cls, + model: str, + effort: str | None, + custom_llm_provider: str, + ) -> str | None: + return cls._validate_effort_for_model(model, effort, custom_llm_provider) + @staticmethod def _model_supports_effort_param(model: str, custom_llm_provider: str) -> bool: """Whether the model accepts ``output_config.effort`` at all. @@ -432,6 +441,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): for level in ("low", "minimal", "medium", "high", "xhigh", "max") ) + @classmethod + def model_supports_effort_param( + cls, + model: str, + custom_llm_provider: str, + ) -> bool: + return cls._model_supports_effort_param(model, custom_llm_provider) + @staticmethod def _model_supports_speed_param(model: str, custom_llm_provider: str | None = None) -> bool: """Whether the model accepts Anthropic's ``speed`` parameter (fast mode). @@ -472,6 +489,16 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) optional_params.pop("speed", None) + @classmethod + def maybe_drop_speed_param( + cls, + model: str, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + drop_params: bool, + custom_llm_provider: str | None = None, + ) -> None: + return cls._maybe_drop_speed_param(model, optional_params, drop_params, custom_llm_provider) + @staticmethod def _raise_invalid_reasoning_effort(model: str, value: object, llm_provider: str) -> NoReturn: """Raise a ``BadRequestError`` for an unrecognised ``reasoning_effort``. @@ -495,6 +522,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): llm_provider=llm_provider, ) + @classmethod + def raise_invalid_reasoning_effort( + cls, + model: str, + value: object, + llm_provider: str, + ) -> NoReturn: + return cls._raise_invalid_reasoning_effort(model, value, llm_provider) + def get_supported_openai_params(self, model: str): params: Final = [ "stream", @@ -923,6 +959,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return returned_tool, mcp_server + def map_tool_helper( + self, + tool: ChatCompletionToolParam, + ) -> tuple[AllAnthropicToolsValues | None, AnthropicMcpServerTool | None]: + return self._map_tool_helper(tool) + def _map_openai_mcp_server_tool(self, tool: OpenAIMcpServerTool) -> AnthropicMcpServerTool: from litellm.types.llms.anthropic import AnthropicMcpServerToolConfiguration @@ -998,6 +1040,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): mcp_servers.append(mcp_server_tool) return anthropic_tools, mcp_servers + def map_tools( + self, + tools: list[ # mutable-ok: mirrors override contract + ChatCompletionToolParam | AllAnthropicToolsValues | dict[str, object] + ], + ) -> tuple[list[AllAnthropicToolsValues], list[AnthropicMcpServerTool]]: # mutable-ok: mirrors override contract + return self._map_tools(tools) + @staticmethod def _rewrite_tool_names_in_messages( messages: list[AllMessageValues], @@ -1249,6 +1299,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): new_stop = new_v return new_stop + def map_stop_sequences( + self, + stop: str | list[str] | None, # mutable-ok: mirrors override contract + ) -> list[str] | None: # mutable-ok: mirrors override contract + return self._map_stop_sequences(stop) + @staticmethod def _map_reasoning_effort( reasoning_effort: REASONING_EFFORT | str | None, @@ -1310,6 +1366,16 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): llm_provider=llm_provider, ) + @classmethod + def map_reasoning_effort( + cls, + reasoning_effort: REASONING_EFFORT | str | None, + model: str, + custom_llm_provider: str, + llm_provider: str = "anthropic", + ) -> AnthropicThinkingParam | None: + return cls._map_reasoning_effort(reasoning_effort, model, custom_llm_provider, llm_provider) + @staticmethod def cap_thinking_budget_to_max_tokens( thinking: AnthropicThinkingParam, max_tokens: int | None @@ -2152,11 +2218,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): json_tool: Final = tool_calls[json_indices[0]] if json_tool.get("function", {}).get("arguments") is None: return None, tool_calls, None - _message: Final = AnthropicConfig._convert_tool_response_to_message(tool_calls=[json_tool]) + _message: Final = AnthropicConfig.convert_tool_response_to_message(tool_calls=[json_tool]) return _message, [], None first_json: Final = tool_calls[json_indices[0]] - json_msg: Final = AnthropicConfig._convert_tool_response_to_message([first_json]) + json_msg: Final = AnthropicConfig.convert_tool_response_to_message([first_json]) extra_content: Final[str | None] = json_msg.content if json_msg is not None else None filtered_tools: Final = [t for i, t in enumerate(tool_calls) if i not in json_indices] return None, filtered_tools, extra_content @@ -2786,6 +2852,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return litellm.Message(content=json_mode_content_str) return None + @classmethod + def convert_tool_response_to_message( + cls, + tool_calls: list[ChatCompletionToolCallChunk], # mutable-ok: mirrors override contract + ) -> LitellmMessage | None: + return cls._convert_tool_response_to_message(tool_calls) + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return AnthropicError( status_code=status_code, diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 69016f5c211..16de0e119c2 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -664,6 +664,18 @@ class AnthropicModelInfo(BaseLLMModelInfo): status_code=400, ) + @classmethod + def apply_sampling_param( + cls, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + model: str, + param: str, + value: object, + drop_params: bool, + output_key: str, + ) -> None: + return cls._apply_sampling_param(optional_params, model, param, value, drop_params, output_key) + @staticmethod def forced_tool_use_unsupported(model: str) -> bool: return AnthropicModelInfo._get_model_capability(model, "supports_forced_tool_use") is False @@ -829,6 +841,15 @@ class AnthropicModelInfo(BaseLLMModelInfo): pass return AnthropicModelInfo._get_model_capability(model, key) is True + @classmethod + def supports_model_capability( + cls, + model: str, + key: str, + custom_llm_provider: str, + ) -> bool: + return cls._supports_model_capability(model, key, custom_llm_provider) + @staticmethod def _is_adaptive_thinking_model(model: str, custom_llm_provider: str) -> bool: """Whether ``model`` uses adaptive thinking (``output_config.effort``). @@ -841,6 +862,14 @@ class AnthropicModelInfo(BaseLLMModelInfo): """ return AnthropicModelInfo._supports_model_capability(model, "supports_adaptive_thinking", custom_llm_provider) + @classmethod + def is_adaptive_thinking_model( + cls, + model: str, + custom_llm_provider: str, + ) -> bool: + return cls._is_adaptive_thinking_model(model, custom_llm_provider) + @staticmethod def _is_always_on_thinking_model(model: str, custom_llm_provider: str) -> bool: """Whether ``model`` always thinks and rejects ``thinking.type=disabled`` diff --git a/litellm/llms/anthropic/completion/transformation.py b/litellm/llms/anthropic/completion/transformation.py index 8edd4354f20..e3ad961ed76 100644 --- a/litellm/llms/anthropic/completion/transformation.py +++ b/litellm/llms/anthropic/completion/transformation.py @@ -163,7 +163,7 @@ class AnthropicTextConfig(BaseConfig): if param == "stream" and value is True: optional_params["stream"] = value if param == "stop" and (isinstance(value, str) or isinstance(value, list)): - _value = litellm.AnthropicConfig()._map_stop_sequences(value) + _value = litellm.AnthropicConfig().map_stop_sequences(value) if _value is not None: optional_params["stop_sequences"] = _value if param == "temperature": @@ -232,6 +232,13 @@ class AnthropicTextConfig(BaseConfig): def _is_anthropic_text_model(model: str) -> bool: return model == "claude-2" or model == "claude-instant-1" + @classmethod + def is_anthropic_text_model( + cls, + model: str, + ) -> bool: + return cls._is_anthropic_text_model(model) + def _get_anthropic_text_prompt_from_messages(self, messages: list[AllMessageValues], model: str) -> str: custom_prompt_dict: Final = litellm.custom_prompt_dict if model in custom_prompt_dict: diff --git a/litellm/llms/anthropic/pass_through/adapters/handler.py b/litellm/llms/anthropic/pass_through/adapters/handler.py index 5bda3437a40..731b39e681f 100644 --- a/litellm/llms/anthropic/pass_through/adapters/handler.py +++ b/litellm/llms/anthropic/pass_through/adapters/handler.py @@ -231,7 +231,7 @@ def _normalize_spec_edits( ) -> list[dict[str, object]] | None: """Return the normalized ``edits`` list, or ``None`` if the polyfill won't run. - Delegates spec-shape normalization to the dispatcher's ``_normalize_spec`` + Delegates spec-shape normalization to the dispatcher's ``normalize_spec`` so the prediction here can't drift from what the dispatcher actually does. """ if not context_management_spec: @@ -241,11 +241,11 @@ def _normalize_spec_edits( return None from litellm.llms.anthropic.pass_through.context_management.dispatcher import ( - _normalize_spec, + normalize_spec, ) try: - return _normalize_spec(context_management_spec) + return normalize_spec(context_management_spec) except Exception: return None diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 6432469366c..2b7f52e0ddc 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -396,7 +396,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): from .transformation import LiteLLMAnthropicMessagesAdapter - usage_dict: UsageDelta = LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage_delta( + usage_dict: UsageDelta = LiteLLMAnthropicMessagesAdapter.translate_openai_usage_to_anthropic_usage_delta( chunk.usage ) if self.applied_edits and "context_management" not in merged_chunk: @@ -1163,7 +1163,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): ( block_type, content_block_start, - ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( + ) = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_chunk_to_anthropic_content_block( choices=chunk.choices ) diff --git a/litellm/llms/anthropic/pass_through/adapters/transformation.py b/litellm/llms/anthropic/pass_through/adapters/transformation.py index 8e0dc83e563..2ee23852894 100644 --- a/litellm/llms/anthropic/pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/pass_through/adapters/transformation.py @@ -1502,7 +1502,7 @@ class LiteLLMAnthropicMessagesAdapter: return cls._first_positive_prompt_tokens_detail_value(usage, ("web_search_requests",)) @classmethod - def _translate_openai_usage_to_anthropic_usage_delta(cls, usage: Usage) -> UsageDelta: + def translate_openai_usage_to_anthropic_usage_delta(cls, usage: Usage) -> UsageDelta: cache_read_input_tokens: Final = cls._get_cache_read_input_tokens(usage) cache_creation_input_tokens: Final = cls._get_cache_creation_input_tokens(usage) web_search_requests: Final = cls._get_web_search_request_count(usage) @@ -1526,13 +1526,17 @@ class LiteLLMAnthropicMessagesAdapter: ) return usage_delta + _translate_openai_usage_to_anthropic_usage_delta = translate_openai_usage_to_anthropic_usage_delta + @classmethod - def _translate_openai_usage_to_anthropic_usage(cls, usage: Usage) -> AnthropicUsage: + def translate_openai_usage_to_anthropic_usage(cls, usage: Usage) -> AnthropicUsage: return cast( AnthropicUsage, - cls._translate_openai_usage_to_anthropic_usage_delta(usage), + cls.translate_openai_usage_to_anthropic_usage_delta(usage), ) + _translate_openai_usage_to_anthropic_usage = translate_openai_usage_to_anthropic_usage + def translate_openai_response_to_anthropic( self, response: ModelResponse, @@ -1576,7 +1580,7 @@ class LiteLLMAnthropicMessagesAdapter: ) # extract usage usage: Final[Usage] = getattr(response, "usage") - message_usage: Final = self._translate_openai_usage_to_anthropic_usage(usage) + message_usage: Final = self.translate_openai_usage_to_anthropic_usage(usage) polyfill_iterations: Final = polyfill_result.iterations_usage if polyfill_result is not None else None anthropic_usage: Final[AnthropicUsage] = ( TypeAdapter(AnthropicUsage).validate_python( @@ -1616,7 +1620,7 @@ class LiteLLMAnthropicMessagesAdapter: return translated_obj - def _translate_streaming_openai_chunk_to_anthropic_content_block( + def translate_streaming_openai_chunk_to_anthropic_content_block( self, choices: list[OpenAIStreamingChoice | StreamingChoices] ) -> tuple[ Literal["text", "tool_use", "thinking"], @@ -1676,6 +1680,10 @@ class LiteLLMAnthropicMessagesAdapter: return "text", TextBlock(type="text", text="") + _translate_streaming_openai_chunk_to_anthropic_content_block = ( + translate_streaming_openai_chunk_to_anthropic_content_block + ) + def _translate_streaming_openai_chunk_to_anthropic( self, choices: list[OpenAIStreamingChoice | StreamingChoices] ) -> tuple[ @@ -1748,7 +1756,7 @@ class LiteLLMAnthropicMessagesAdapter: else None ) if litellm_usage_chunk is not None: - usage_delta = self._translate_openai_usage_to_anthropic_usage_delta(litellm_usage_chunk) + usage_delta = self.translate_openai_usage_to_anthropic_usage_delta(litellm_usage_chunk) else: usage_delta = UsageDelta(input_tokens=0, output_tokens=0) message_block: Final = MessageBlockDelta( diff --git a/litellm/llms/anthropic/pass_through/context_management/dispatcher.py b/litellm/llms/anthropic/pass_through/context_management/dispatcher.py index ad33e5e0592..09fbaa89237 100644 --- a/litellm/llms/anthropic/pass_through/context_management/dispatcher.py +++ b/litellm/llms/anthropic/pass_through/context_management/dispatcher.py @@ -33,7 +33,7 @@ def _edits_from(normalized: dict[str, object] | None) -> list[dict[str, object]] return [edit for edit in edits if isinstance(edit, dict)] -def _normalize_spec( +def normalize_spec( spec: dict[str, object] | list[dict[str, object]] | None, ) -> list[dict[str, object]] | None: """Accept Anthropic-native dict form or OpenAI list form; return edits list.""" @@ -46,6 +46,9 @@ def _normalize_spec( return _edits_from(spec) +_normalize_spec = normalize_spec + + def _wrap_editor_return( raw: EditorResult, *, @@ -87,7 +90,7 @@ async def apply_context_management( worker thread so their token counts stay off the event loop; ``inspect.iscoroutinefunction`` decides how each editor is invoked. """ - edits: Final = _normalize_spec(context_management_spec) + edits: Final = normalize_spec(context_management_spec) if not edits: return PolyfillResult(messages=messages, system=system, applied_edits=[]) diff --git a/litellm/llms/anthropic/pass_through/messages/response_cache.py b/litellm/llms/anthropic/pass_through/messages/response_cache.py index f35c4c80f2b..4c5f47bc650 100644 --- a/litellm/llms/anthropic/pass_through/messages/response_cache.py +++ b/litellm/llms/anthropic/pass_through/messages/response_cache.py @@ -9,9 +9,9 @@ from litellm.caching.caching_handler import create_cache_write_task, is_response from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( AnthropicMessagesStreamingResponse, BaseAnthropicMessagesStreamingIterator, - _is_message_stop_chunk, - _is_provider_error_chunk, aclose_if_supported, + is_message_stop_chunk, + is_provider_error_chunk, ) from litellm.types.caching import CACHED_STREAM_EVENTS_KEY @@ -90,7 +90,7 @@ class AnthropicMessagesStreamCacheWriter: if self.persisted or cache is None: return collected_stream: Final = b"".join(self.collected_chunks) - if not _is_message_stop_chunk(collected_stream) or _is_provider_error_chunk(collected_stream): + if not is_message_stop_chunk(collected_stream) or is_provider_error_chunk(collected_stream): return self.persisted = True diff --git a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py index 150b1a730ef..25e0f93b04c 100644 --- a/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/messages/streaming_iterator.py @@ -30,7 +30,7 @@ _UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdl _DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains -def _is_message_stop_chunk(chunk: object) -> bool: +def is_message_stop_chunk(chunk: object) -> bool: if isinstance(chunk, dict): return chunk.get("type") == "message_stop" if isinstance(chunk, (bytes, bytearray)): @@ -38,6 +38,9 @@ def _is_message_stop_chunk(chunk: object) -> bool: return False +_is_message_stop_chunk = is_message_stop_chunk + + def is_anthropic_ping_chunk(chunk: object) -> bool: """ Whether a chunk is made only of whole ``ping`` keepalive frames. A ping @@ -133,10 +136,13 @@ def _anthropic_error_body(chunk: object) -> Mapping[str, object] | None: return error_body if isinstance(error_body, dict) else None -def _is_provider_error_chunk(chunk: object) -> bool: +def is_provider_error_chunk(chunk: object) -> bool: return _anthropic_error_body(chunk) is not None +_is_provider_error_chunk = is_provider_error_chunk + + def parse_anthropic_error_event(chunk: object) -> tuple[str, str, int] | None: """ Extract ``(error_type, message, http_status_code)`` from an Anthropic SSE @@ -161,7 +167,7 @@ def parse_anthropic_error_event(chunk: object) -> tuple[str, str, int] | None: def _is_terminal_stream_chunk(chunk: object) -> bool: - return _is_message_stop_chunk(chunk) or _is_provider_error_chunk(chunk) + return is_message_stop_chunk(chunk) or is_provider_error_chunk(chunk) def _try_claim_detached_drain_slot() -> bool: diff --git a/litellm/llms/anthropic/pass_through/messages/transformation.py b/litellm/llms/anthropic/pass_through/messages/transformation.py index b35c35f1411..f1eb30b63c1 100644 --- a/litellm/llms/anthropic/pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/pass_through/messages/transformation.py @@ -62,10 +62,13 @@ DROP_UNFITTING_REASONING_EFFORT_WARNING: Final = ( ) -def _messages_carry_output_config(messages: Sequence[object]) -> bool: +def messages_carry_output_config(messages: Sequence[object]) -> bool: return any(isinstance(message, Mapping) and "output_config" in message for message in messages) +_messages_carry_output_config = messages_carry_output_config + + class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): _workload_identity_eligible: ClassVar[bool] = True @@ -395,7 +398,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): return try: - mapped_thinking: Final = AnthropicConfig._map_reasoning_effort( + mapped_thinking: Final = AnthropicConfig.map_reasoning_effort( reasoning_effort=reasoning_effort, model=model, custom_llm_provider=custom_llm_provider, @@ -414,7 +417,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): return optional_params.setdefault("thinking", fitted_thinking) - if AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider): + if AnthropicModelInfo.is_adaptive_thinking_model(model, custom_llm_provider): mapped_effort: Final = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(reasoning_effort) if mapped_effort is None: raise AnthropicError( @@ -425,7 +428,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): ), status_code=400, ) - gate_error: Final = AnthropicConfig._validate_effort_for_model(model, mapped_effort, custom_llm_provider) + gate_error: Final = AnthropicConfig.validate_effort_for_model(model, mapped_effort, custom_llm_provider) if gate_error is not None: raise AnthropicError(message=gate_error, status_code=400) existing_output_config = optional_params.get("output_config") @@ -479,7 +482,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): from litellm.exceptions import BadRequestError as _BadRequestError from litellm.llms.anthropic.chat.transformation import AnthropicConfig - if AnthropicConfig._is_adaptive_thinking_model(model, custom_llm_provider): + if AnthropicConfig.is_adaptive_thinking_model(model, custom_llm_provider): return output_config: Final = optional_params.get("output_config") @@ -494,20 +497,20 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): # reject. Effort-only requests pass through so provider subclasses (bedrock/vertex) keep # owning level clamping; an adaptive request only stays here when its effort level is one # the model supports, otherwise it falls through to the legacy budget translation below. - if AnthropicConfig._model_supports_effort_param(model, custom_llm_provider) and ( + if AnthropicConfig.model_supports_effort_param(model, custom_llm_provider) and ( not adaptive_thinking - or AnthropicConfig._validate_effort_for_model(model, effort, custom_llm_provider) is None + or AnthropicConfig.validate_effort_for_model(model, effort, custom_llm_provider) is None ): if adaptive_thinking: optional_params.pop("thinking", None) return - supports_thinking: Final = AnthropicModelInfo._supports_model_capability( + supports_thinking: Final = AnthropicModelInfo.supports_model_capability( model, "supports_reasoning", custom_llm_provider ) try: legacy_thinking: Final = ( - AnthropicConfig._map_reasoning_effort( + AnthropicConfig.map_reasoning_effort( reasoning_effort=effort or "medium", model=model, custom_llm_provider=custom_llm_provider, @@ -554,7 +557,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): Adaptive models (4.6+) own this natively and are left untouched. """ - if AnthropicModelInfo._is_adaptive_thinking_model(model, custom_llm_provider): + if AnthropicModelInfo.is_adaptive_thinking_model(model, custom_llm_provider): return temperature: Final = optional_params.get("temperature") if temperature is None or temperature == 1: @@ -770,7 +773,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): if optional_params.get("speed") == "fast": beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.FAST_MODE_2026_02_01.value) - if _messages_carry_output_config(messages): + if messages_carry_output_config(messages): beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.PER_TURN_CONTROL_2026_07_01.value) tools: Final = optional_params.get("tools") diff --git a/litellm/llms/anthropic/pass_through/messages/utils.py b/litellm/llms/anthropic/pass_through/messages/utils.py index fe8ac2cd7a2..dfab0af8eaa 100644 --- a/litellm/llms/anthropic/pass_through/messages/utils.py +++ b/litellm/llms/anthropic/pass_through/messages/utils.py @@ -158,7 +158,7 @@ class AnthropicMessagesRequestUtils: from litellm.llms.anthropic.chat.transformation import AnthropicConfig from litellm.llms.anthropic.common_utils import AnthropicModelInfo - AnthropicConfig._maybe_drop_speed_param( + AnthropicConfig.maybe_drop_speed_param( model=model, optional_params=filtered_params, drop_params=drop_params, diff --git a/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py index 88503908c3c..642f1857058 100644 --- a/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/pass_through/responses_adapters/transformation.py @@ -75,7 +75,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: from litellm.responses.utils import ResponseAPILoggingUtils chat_usage = ResponseAPILoggingUtils.transform_response_api_usage_to_chat_usage(raw_usage) - return LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage(chat_usage) + return LiteLLMAnthropicMessagesAdapter.translate_openai_usage_to_anthropic_usage(chat_usage) # ------------------------------------------------------------------ # # Request translation: Anthropic -> Responses API # diff --git a/litellm/llms/azure/audio_transcriptions.py b/litellm/llms/azure/audio_transcriptions.py index 564ec94ba6b..99ebd2875a7 100644 --- a/litellm/llms/azure/audio_transcriptions.py +++ b/litellm/llms/azure/audio_transcriptions.py @@ -78,7 +78,7 @@ class AzureAudioTranscription(AzureChatCompletion): api_key=azure_client.api_key, additional_args={ "headers": {"Authorization": f"Bearer {azure_client.api_key}"}, - "api_base": azure_client._base_url._uri_reference, + "api_base": azure_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "atranscription": True, "complete_input_dict": data, }, @@ -149,7 +149,7 @@ class AzureAudioTranscription(AzureChatCompletion): api_key=async_azure_client.api_key, additional_args={ "headers": {"Authorization": f"Bearer {async_azure_client.api_key}"}, - "api_base": async_azure_client._base_url._uri_reference, + "api_base": async_azure_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "atranscription": True, "complete_input_dict": data, }, @@ -175,7 +175,7 @@ class AzureAudioTranscription(AzureChatCompletion): api_key=api_key, additional_args={ "headers": {"Authorization": f"Bearer {async_azure_client.api_key}"}, - "api_base": async_azure_client._base_url._uri_reference, + "api_base": async_azure_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "atranscription": True, "complete_input_dict": data, }, diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 64441fb0a43..d396e1f4867 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -112,7 +112,7 @@ class AzureOpenAIAssistantsAPIConfig: return optional_params -def _check_dynamic_azure_params( +def check_dynamic_azure_params( azure_client_params: dict, azure_client: AzureOpenAI | AsyncAzureOpenAI | None, ) -> bool: @@ -127,12 +127,15 @@ def _check_dynamic_azure_params( dynamic_params: Final = ["api_version"] for k, v in azure_client_params.items(): if k in dynamic_params and k == "api_version": - if v is not None and v != azure_client._custom_query["api-version"]: + if v is not None and v != azure_client._custom_query["api-version"]: # pyright: ignore[reportPrivateUsage] # SDK query internals return True return False +_check_dynamic_azure_params = check_dynamic_azure_params + + class AzureChatCompletion(BaseAzureLLM, BaseLLM): def __init__(self) -> None: super().__init__() diff --git a/litellm/llms/azure/chat/gpt_5_transformation.py b/litellm/llms/azure/chat/gpt_5_transformation.py index 3189f5b57ac..82150e415aa 100644 --- a/litellm/llms/azure/chat/gpt_5_transformation.py +++ b/litellm/llms/azure/chat/gpt_5_transformation.py @@ -6,7 +6,7 @@ import litellm from litellm.exceptions import UnsupportedParamsError from litellm.llms.openai.chat.gpt_5_transformation import ( OpenAIGPT5Config, - _get_effort_level, + get_effort_level, is_gpt_reasoning_series_name, ) from litellm.types.llms.openai import AllMessageValues @@ -76,7 +76,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): api_version: str = "", ) -> dict: reasoning_effort_value = non_default_params.get("reasoning_effort") or optional_params.get("reasoning_effort") - effective_effort: Final = _get_effort_level(reasoning_effort_value) + effective_effort: Final = get_effort_level(reasoning_effort_value) # gpt-5.1/5.2/5.4 support reasoning_effort='none', but other gpt-5 models don't # See: https://learn.microsoft.com/en-us/azure/ai-foundry/openai/how-to/reasoning @@ -86,9 +86,9 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): if litellm.drop_params is True or (drop_params is not None and drop_params is True): non_default_params = non_default_params.copy() optional_params = optional_params.copy() - if _get_effort_level(non_default_params.get("reasoning_effort")) == "none": + if get_effort_level(non_default_params.get("reasoning_effort")) == "none": non_default_params.pop("reasoning_effort") - if _get_effort_level(optional_params.get("reasoning_effort")) == "none": + if get_effort_level(optional_params.get("reasoning_effort")) == "none": optional_params.pop("reasoning_effort") else: raise UnsupportedParamsError( @@ -111,7 +111,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): ) # Only drop reasoning_effort='none' for models that don't support it - result_effort: Final = _get_effort_level(result.get("reasoning_effort")) + result_effort: Final = get_effort_level(result.get("reasoning_effort")) if result_effort == "none" and not supports_none: result.pop("reasoning_effort") diff --git a/litellm/llms/azure/common_utils.py b/litellm/llms/azure/common_utils.py index ea6e9bdabad..c640405c269 100644 --- a/litellm/llms/azure/common_utils.py +++ b/litellm/llms/azure/common_utils.py @@ -388,7 +388,7 @@ def get_azure_ad_token( # try to get DefaultAzureCredential provider ######################################################### if azure_ad_token_provider is None and azure_ad_token is None: - azure_ad_token_provider = BaseAzureLLM._try_get_default_azure_credential_provider( + azure_ad_token_provider = BaseAzureLLM.try_get_default_azure_credential_provider( scope=scope, ) @@ -477,6 +477,13 @@ class BaseAzureLLM(BaseOpenAILLM): verbose_logger.debug("DefaultAzureCredential failed: %s", e) return None + @classmethod + def try_get_default_azure_credential_provider( + cls, + scope: str, + ) -> Callable[[], str] | None: + return cls._try_get_default_azure_credential_provider(scope) + def get_azure_openai_client( self, api_key: str | None, @@ -511,10 +518,12 @@ class BaseAzureLLM(BaseOpenAILLM): if ( api_version is not None and isinstance(client, (AzureOpenAI, AsyncAzureOpenAI)) - and isinstance(client._custom_query, dict) + and isinstance(client._custom_query, dict) # pyright: ignore[reportPrivateUsage] # SDK query internals ): # set api_version to version passed by user - client._custom_query.setdefault("api-version", api_version) + client._custom_query.setdefault( # pyright: ignore[reportPrivateUsage] # SDK query internals + "api-version", api_version + ) self.set_cached_openai_client( openai_client=client, client_initialization_params=client_initialization_params, @@ -777,6 +786,14 @@ class BaseAzureLLM(BaseOpenAILLM): return headers + @classmethod + def base_validate_azure_environment( + cls, + headers: dict[str, str], # mutable-ok: mirrors override contract + litellm_params: GenericLiteLLMParams | None, + ) -> dict[str, str]: # mutable-ok: mirrors override contract + return cls._base_validate_azure_environment(headers, litellm_params) + @staticmethod def _get_base_azure_url( api_base: str | None, @@ -829,6 +846,16 @@ class BaseAzureLLM(BaseOpenAILLM): return str(final_url) + @classmethod + def get_base_azure_url( + cls, + api_base: str | None, + litellm_params: GenericLiteLLMParams | Mapping[str, object] | None, + route: Literal["/openai/responses", "/openai/vector_stores"] | str, + default_api_version: str | Literal["latest", "preview"] | None = None, + ) -> str: + return cls._get_base_azure_url(api_base, litellm_params, route, default_api_version) + @staticmethod def get_azure_v1_image_url(api_base: str, api_version: str | None, route: str) -> str | None: """ diff --git a/litellm/llms/azure/completion/handler.py b/litellm/llms/azure/completion/handler.py index 476f6536ffa..34e51eaf1ec 100644 --- a/litellm/llms/azure/completion/handler.py +++ b/litellm/llms/azure/completion/handler.py @@ -225,7 +225,7 @@ class AzureTextCompletion(BaseAzureLLM): api_key=azure_client.api_key, additional_args={ "headers": {"Authorization": f"Bearer {azure_client.api_key}"}, - "api_base": azure_client._base_url._uri_reference, + "api_base": azure_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "acompletion": True, "complete_input_dict": data, }, @@ -284,7 +284,7 @@ class AzureTextCompletion(BaseAzureLLM): api_key=azure_client.api_key, additional_args={ "headers": {"Authorization": f"Bearer {azure_client.api_key}"}, - "api_base": azure_client._base_url._uri_reference, + "api_base": azure_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "acompletion": True, "complete_input_dict": data, }, @@ -334,7 +334,7 @@ class AzureTextCompletion(BaseAzureLLM): api_key=azure_client.api_key, additional_args={ "headers": {"Authorization": f"Bearer {azure_client.api_key}"}, - "api_base": azure_client._base_url._uri_reference, + "api_base": azure_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "acompletion": True, "complete_input_dict": data, }, diff --git a/litellm/llms/azure/containers/transformation.py b/litellm/llms/azure/containers/transformation.py index 11a2ded91ad..5031f38bf8f 100644 --- a/litellm/llms/azure/containers/transformation.py +++ b/litellm/llms/azure/containers/transformation.py @@ -29,7 +29,7 @@ class AzureContainerConfig(OpenAIContainerConfig): headers: dict, api_key: str | None = None, ) -> dict: - return BaseAzureLLM._base_validate_azure_environment( + return BaseAzureLLM.base_validate_azure_environment( headers=headers, litellm_params=GenericLiteLLMParams(api_key=api_key), ) @@ -75,7 +75,7 @@ class AzureContainerConfig(OpenAIContainerConfig): api_version_from_base: Final = self._extract_api_version(api_base) if api_version_from_base: effective_params["api_version"] = api_version_from_base - return BaseAzureLLM._get_base_azure_url( + return BaseAzureLLM.get_base_azure_url( api_base=self._normalize_api_base(api_base), litellm_params=effective_params, route="/openai/containers", diff --git a/litellm/llms/azure/fine_tuning/handler.py b/litellm/llms/azure/fine_tuning/handler.py index 36c4fae04c7..0baf3ccc598 100644 --- a/litellm/llms/azure/fine_tuning/handler.py +++ b/litellm/llms/azure/fine_tuning/handler.py @@ -8,7 +8,7 @@ from litellm._logging import verbose_logger from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.openai.fine_tuning.handler import ( OpenAIFineTuningAPI, - _litellm_fine_tuning_job_from_response, + litellm_fine_tuning_job_from_response, ) from litellm.types.utils import LiteLLMFineTuningJob @@ -37,7 +37,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response: Final = await openai_client.fine_tuning.jobs.create(**create_fine_tuning_job_data) - return _litellm_fine_tuning_job_from_response(response, is_azure=True) + return litellm_fine_tuning_job_from_response(response, is_azure=True) async def acancel_fine_tuning_job( self, @@ -45,7 +45,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response: Final = await openai_client.fine_tuning.jobs.cancel(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response, is_azure=True) + return litellm_fine_tuning_job_from_response(response, is_azure=True) async def aretrieve_fine_tuning_job( self, @@ -53,7 +53,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response: Final = await openai_client.fine_tuning.jobs.retrieve(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response, is_azure=True) + return litellm_fine_tuning_job_from_response(response, is_azure=True) def create_fine_tuning_job( self, @@ -96,7 +96,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): verbose_logger.debug("creating fine tuning job, args= %s", create_fine_tuning_job_data) response: Final = cast(OpenAI, openai_client).fine_tuning.jobs.create(**create_fine_tuning_job_data) - return _litellm_fine_tuning_job_from_response(response, is_azure=True) + return litellm_fine_tuning_job_from_response(response, is_azure=True) def cancel_fine_tuning_job( self, @@ -136,7 +136,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): ) response: Final = cast(OpenAI, openai_client).fine_tuning.jobs.cancel(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response, is_azure=True) + return litellm_fine_tuning_job_from_response(response, is_azure=True) def retrieve_fine_tuning_job( self, @@ -176,7 +176,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): ) response: Final = cast(OpenAI, openai_client).fine_tuning.jobs.retrieve(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response, is_azure=True) + return litellm_fine_tuning_job_from_response(response, is_azure=True) def get_openai_client( self, diff --git a/litellm/llms/azure/image_edit/transformation.py b/litellm/llms/azure/image_edit/transformation.py index 05defa6d0ff..dc963d7b9f8 100644 --- a/litellm/llms/azure/image_edit/transformation.py +++ b/litellm/llms/azure/image_edit/transformation.py @@ -65,7 +65,7 @@ class AzureImageEditConfig(OpenAIImageEditConfig): params: Final = GenericLiteLLMParams(**(litellm_params or {})) if api_key is not None and params.api_key is None: params.api_key = api_key - return BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=params) + return BaseAzureLLM.base_validate_azure_environment(headers=headers, litellm_params=params) def get_complete_url( self, diff --git a/litellm/llms/azure/passthrough/transformation.py b/litellm/llms/azure/passthrough/transformation.py index ff83151626a..501122783eb 100644 --- a/litellm/llms/azure/passthrough/transformation.py +++ b/litellm/llms/azure/passthrough/transformation.py @@ -128,7 +128,7 @@ class AzurePassthroughConfig(BasePassthroughConfig): caller_api_version: Final = request_query_params.get("api-version") if request_query_params else None relay_base: Final = without_api_version(base_target_url) if caller_api_version else base_target_url - complete_url: Final = BaseAzureLLM._get_base_azure_url( + complete_url: Final = BaseAzureLLM.get_base_azure_url( api_base=relay_base, litellm_params=MappingProxyType( {**litellm_params, "api_version": caller_api_version or litellm_params.get("api_version")} @@ -150,7 +150,7 @@ class AzurePassthroughConfig(BasePassthroughConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - return BaseAzureLLM._base_validate_azure_environment( + return BaseAzureLLM.base_validate_azure_environment( headers=headers, litellm_params=GenericLiteLLMParams.model_validate({**litellm_params, "api_key": api_key}), ) diff --git a/litellm/llms/azure/realtime/handler.py b/litellm/llms/azure/realtime/handler.py index df74975ad0f..79dd7f098c8 100644 --- a/litellm/llms/azure/realtime/handler.py +++ b/litellm/llms/azure/realtime/handler.py @@ -130,6 +130,16 @@ class AzureOpenAIRealtime(AzureChatCompletion): qs: Final = "&".join(query_parts) return f"{api_base}{path}?{qs}" if qs else f"{api_base}{path}" + def construct_url( + self, + api_base: str, + model: str, + api_version: str | None, + realtime_protocol: str | None = None, + query_params: RealtimeQueryParams | None = None, + ) -> str: + return self._construct_url(api_base, model, api_version, realtime_protocol, query_params) + async def async_realtime( self, model: str, diff --git a/litellm/llms/azure/responses/transformation.py b/litellm/llms/azure/responses/transformation.py index 2a82b42df7b..75f47b9229c 100644 --- a/litellm/llms/azure/responses/transformation.py +++ b/litellm/llms/azure/responses/transformation.py @@ -32,7 +32,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): @staticmethod def _supports_reasoning_effort_none(model: str) -> bool: - return AzureOpenAIGPT5Config._supports_reasoning_effort_level(model, "none") + return AzureOpenAIGPT5Config.supports_reasoning_effort_level(model, "none") @staticmethod def _effort_resolves_to_none(model: str, effort: str | None) -> bool: @@ -46,7 +46,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): return [param for param in base_supported_params if param not in self.AZURE_UNSUPPORTED_PARAMS] def validate_environment(self, headers: dict, model: str, litellm_params: GenericLiteLLMParams | None) -> dict: - return BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params) + return BaseAzureLLM.base_validate_azure_environment(headers=headers, litellm_params=litellm_params) def get_stripped_model_name(self, model: str) -> str: return model.replace("responses/", "").replace("o_series/", "").replace("azure_ai/", "") @@ -169,7 +169,7 @@ class AzureOpenAIResponsesAPIConfig(OpenAIResponsesAPIConfig): """ from litellm.constants import AZURE_DEFAULT_RESPONSES_API_VERSION - return BaseAzureLLM._get_base_azure_url( + return BaseAzureLLM.get_base_azure_url( api_base=api_base, litellm_params=litellm_params, route="/openai/responses", diff --git a/litellm/llms/azure/vector_stores/transformation.py b/litellm/llms/azure/vector_stores/transformation.py index 7e80ee46046..9510c021c7d 100644 --- a/litellm/llms/azure/vector_stores/transformation.py +++ b/litellm/llms/azure/vector_stores/transformation.py @@ -9,11 +9,11 @@ class AzureOpenAIVectorStoreConfig(OpenAIVectorStoreConfig): api_base: str | None, litellm_params: dict, ) -> str: - return BaseAzureLLM._get_base_azure_url( + return BaseAzureLLM.get_base_azure_url( api_base=api_base, litellm_params=litellm_params, route="/openai/vector_stores", ) def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: - return BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params) + return BaseAzureLLM.base_validate_azure_environment(headers=headers, litellm_params=litellm_params) diff --git a/litellm/llms/azure/videos/transformation.py b/litellm/llms/azure/videos/transformation.py index daca765a1d2..67dff05fb08 100644 --- a/litellm/llms/azure/videos/transformation.py +++ b/litellm/llms/azure/videos/transformation.py @@ -72,7 +72,7 @@ class AzureVideoConfig(OpenAIVideoConfig): # Use the base Azure validation method which properly handles: # 1. Credentials from litellm_credential_name via litellm_params # 2. Sets the correct "api-key" header (not "Authorization: Bearer") - return BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params) + return BaseAzureLLM.base_validate_azure_environment(headers=headers, litellm_params=litellm_params) def get_complete_url( self, @@ -83,7 +83,7 @@ class AzureVideoConfig(OpenAIVideoConfig): """ Constructs a complete URL for the API request. """ - return BaseAzureLLM._get_base_azure_url( + return BaseAzureLLM.get_base_azure_url( api_base=api_base, litellm_params=litellm_params, route="/openai/v1/videos", diff --git a/litellm/llms/azure_ai/agents/handler.py b/litellm/llms/azure_ai/agents/handler.py index 3f0eeef9c7a..0b0548dd3a2 100644 --- a/litellm/llms/azure_ai/agents/handler.py +++ b/litellm/llms/azure_ai/agents/handler.py @@ -279,7 +279,7 @@ class AzureAIAgentsHandler: headers["Authorization"] = f"Bearer {api_key}" api_version: Final = optional_params.get("api_version", self.config.DEFAULT_API_VERSION) - agent_id: Final = self.config._get_agent_id(model, optional_params) + agent_id: Final = self.config.get_agent_id(model, optional_params) thread_id: Final = optional_params.get("thread_id") api_base = api_base.rstrip("/") @@ -313,10 +313,10 @@ class AzureAIAgentsHandler: headers: dict | None = None, ) -> ModelResponse: """Execute synchronous completion using Azure Agent Service.""" - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client if client is None: - client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) ( headers, diff --git a/litellm/llms/azure_ai/agents/transformation.py b/litellm/llms/azure_ai/agents/transformation.py index 0b6ea3ba717..18f8f7916b7 100644 --- a/litellm/llms/azure_ai/agents/transformation.py +++ b/litellm/llms/azure_ai/agents/transformation.py @@ -184,6 +184,13 @@ class AzureAIAgentsConfig(BaseConfig): # Extract from model name using the static method return self.get_agent_id_from_model(model) + def get_agent_id( + self, + model: str, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> str: + return self._get_agent_id(model, optional_params) + def transform_request( self, model: str, diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py b/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py index e96a2f20bbb..d7455439151 100644 --- a/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py +++ b/litellm/llms/azure_ai/anthropic/count_tokens/transformation.py @@ -56,7 +56,7 @@ class AzureAIAnthropicCountTokensConfig(AnthropicCountTokensConfig): litellm_params_obj: Final = GenericLiteLLMParams.model_validate(litellm_params) # Get Azure auth headers (api-key or Authorization) - azure_headers = BaseAzureLLM._base_validate_azure_environment(headers={}, litellm_params=litellm_params_obj) + azure_headers = BaseAzureLLM.base_validate_azure_environment(headers={}, litellm_params=litellm_params_obj) # Merge Azure auth headers headers.update(azure_headers) diff --git a/litellm/llms/azure_ai/anthropic/handler.py b/litellm/llms/azure_ai/anthropic/handler.py index a132df62b08..6065c56d083 100644 --- a/litellm/llms/azure_ai/anthropic/handler.py +++ b/litellm/llms/azure_ai/anthropic/handler.py @@ -177,9 +177,9 @@ class AzureAnthropicChatCompletion(AnthropicChatCompletion): else: if client is None or not isinstance(client, HTTPHandler): - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client - client = _get_httpx_client(params={"timeout": timeout}) + client = get_httpx_client(params={"timeout": timeout}) else: client = client diff --git a/litellm/llms/azure_ai/anthropic/messages_transformation.py b/litellm/llms/azure_ai/anthropic/messages_transformation.py index 3d8b574e8c0..6803bf7e551 100644 --- a/litellm/llms/azure_ai/anthropic/messages_transformation.py +++ b/litellm/llms/azure_ai/anthropic/messages_transformation.py @@ -50,7 +50,7 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): litellm_params_obj.api_key = api_key # Use Azure authentication logic - headers = BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params_obj) + headers = BaseAzureLLM.base_validate_azure_environment(headers=headers, litellm_params=litellm_params_obj) # Azure Anthropic uses x-api-key header (not api-key) # Convert api-key to x-api-key if present diff --git a/litellm/llms/azure_ai/anthropic/transformation.py b/litellm/llms/azure_ai/anthropic/transformation.py index 864d2134a84..18a4eeb7c0a 100644 --- a/litellm/llms/azure_ai/anthropic/transformation.py +++ b/litellm/llms/azure_ai/anthropic/transformation.py @@ -73,7 +73,7 @@ class AzureAnthropicConfig(AnthropicConfig): litellm_params_obj.api_key = api_key # Use Azure authentication logic - headers = BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params_obj) + headers = BaseAzureLLM.base_validate_azure_environment(headers=headers, litellm_params=litellm_params_obj) # Get tools and other anthropic-specific setup tools: Final = optional_params.get("tools") diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index fc5e76c9295..a8e6320bd1e 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -117,7 +117,7 @@ class AzureAIStudioConfig(OpenAIConfig): if "grok" in model: # Reuse Xai method for Grok model xai_config: Final = XAIChatConfig() - return xai_config._supports_stop_reason(model) + return xai_config.supports_stop_reason(model) return True def validate_environment( @@ -138,7 +138,7 @@ class AzureAIStudioConfig(OpenAIConfig): else: # No api_key provided — fall back to Azure AD token-based auth litellm_params_obj = GenericLiteLLMParams(**(litellm_params if isinstance(litellm_params, dict) else {})) - headers = BaseAzureLLM._base_validate_azure_environment(headers=headers, litellm_params=litellm_params_obj) + headers = BaseAzureLLM.base_validate_azure_environment(headers=headers, litellm_params=litellm_params_obj) headers["Content-Type"] = "application/json" @@ -281,6 +281,15 @@ class AzureAIStudioConfig(OpenAIConfig): custom_llm_provider = "azure" return api_base, dynamic_api_key, custom_llm_provider + def get_openai_compatible_provider_info( + self, + model: str, + api_base: str | None, + api_key: str | None, + custom_llm_provider: str, + ) -> tuple[str | None, str | None, str]: + return self._get_openai_compatible_provider_info(model, api_base, api_key, custom_llm_provider) + def transform_request( self, model: str, diff --git a/litellm/llms/azure_ai/embed/cohere_transformation.py b/litellm/llms/azure_ai/embed/cohere_transformation.py index 453e48d6819..82e855d8706 100644 --- a/litellm/llms/azure_ai/embed/cohere_transformation.py +++ b/litellm/llms/azure_ai/embed/cohere_transformation.py @@ -67,6 +67,14 @@ class AzureAICohereConfig: return image_embeddings_request, v1_embeddings_request, image_embedding_idx + def transform_request( + self, + input: list[str], # mutable-ok: mirrors override contract + optional_params: dict[str, object], # mutable-ok: mirrors override contract + model: str, + ) -> tuple[ImageEmbeddingRequest, EmbeddingCreateParams, list[int]]: # mutable-ok: mirrors override contract + return self._transform_request(input, optional_params, model) + def _transform_response(self, response: EmbeddingResponse) -> EmbeddingResponse: additional_headers: Final[dict | None] = response.hidden_params.get("additional_headers") if additional_headers: @@ -84,3 +92,9 @@ class AzureAICohereConfig: response.model = self._map_azure_model_group(base_model) return response + + def transform_response( + self, + response: EmbeddingResponse, + ) -> EmbeddingResponse: + return self._transform_response(response) diff --git a/litellm/llms/azure_ai/embed/handler.py b/litellm/llms/azure_ai/embed/handler.py index 8037be8abeb..57e6ca51ade 100644 --- a/litellm/llms/azure_ai/embed/handler.py +++ b/litellm/llms/azure_ai/embed/handler.py @@ -58,7 +58,7 @@ class AzureAIEmbedding(OpenAIChatCompletion): elif text_embedding_responses is not None: model_response.data = text_embedding_responses - response: Final = AzureAICohereConfig()._transform_response(response=model_response) + response: Final = AzureAICohereConfig().transform_response(response=model_response) return response @@ -158,7 +158,7 @@ class AzureAIEmbedding(OpenAIChatCompletion): image_embeddings_request, v1_embeddings_request, image_embeddings_idx, - ) = AzureAICohereConfig()._transform_request(input=input, optional_params=optional_params, model=model) + ) = AzureAICohereConfig().transform_request(input=input, optional_params=optional_params, model=model) image_embedding_responses: list | None = None text_embedding_responses: list | None = None @@ -246,7 +246,7 @@ class AzureAIEmbedding(OpenAIChatCompletion): image_embeddings_request, v1_embeddings_request, image_embeddings_idx, - ) = AzureAICohereConfig()._transform_request(input=input, optional_params=optional_params, model=model) + ) = AzureAICohereConfig().transform_request(input=input, optional_params=optional_params, model=model) image_embedding_responses: list | None = None text_embedding_responses: list | None = None diff --git a/litellm/llms/base_llm/base_model_iterator.py b/litellm/llms/base_llm/base_model_iterator.py index c6bd4e660c6..9cb5c2c55dd 100644 --- a/litellm/llms/base_llm/base_model_iterator.py +++ b/litellm/llms/base_llm/base_model_iterator.py @@ -108,6 +108,13 @@ class BaseModelResponseIterator: stripped_json_chunk = None return stripped_json_chunk + @classmethod + def string_to_dict_parser( + cls, + str_line: str, + ) -> dict[str, object] | None: # mutable-ok: mirrors override contract + return cls._string_to_dict_parser(str_line) + def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream: # chunk is a str at this point stripped_json_chunk: Final = BaseModelResponseIterator._string_to_dict_parser(str_line=str_line) diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index 9f690b9825a..ba4dad40d13 100644 --- a/litellm/llms/base_llm/base_utils.py +++ b/litellm/llms/base_llm/base_utils.py @@ -8,7 +8,7 @@ from abc import ABC, abstractmethod from collections.abc import Mapping from typing import Any, Final -from openai.lib import _parsing, _pydantic +from openai.lib import _parsing, _pydantic # pyright: ignore[reportPrivateUsage] # SDK parser internals from pydantic import BaseModel from litellm._logging import verbose_logger @@ -128,7 +128,7 @@ class BaseLLMModelInfo(ABC): return None -def _convert_tool_response_to_message( +def convert_tool_response_to_message( tool_calls: list[ChatCompletionToolCallChunk], ) -> Message | None: """ @@ -154,6 +154,9 @@ def _convert_tool_response_to_message( return None +_convert_tool_response_to_message = convert_tool_response_to_message + + def _dict_to_response_format_helper(response_format: dict, ref_template: str | None = None) -> dict: if ref_template is not None and response_format.get("type") == "json_schema": # Deep copy to avoid modifying original @@ -211,7 +214,7 @@ def type_to_response_format_param( # type checkers don't narrow the negation of a `TypeGuard` as it isn't # a safe default behaviour but we know that at this point the `response_format` # can only be a `type` - if not _parsing._completions.is_basemodel_type(response_format): + if not _parsing._completions.is_basemodel_type(response_format): # pyright: ignore[reportPrivateUsage] # SDK parser internals raise TypeError(f"Unsupported response_format type - {response_format}") if ref_template is not None: diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 3885e4ad38a..50825720331 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -146,6 +146,13 @@ class BaseConfig(ABC): ] return optional_params + def add_tools_to_optional_params( + self, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + tools: list[ChatCompletionToolParam], # mutable-ok: mirrors override contract + ) -> dict[str, object]: # mutable-ok: mirrors override contract + return self._add_tools_to_optional_params(optional_params, tools) + def translate_developer_role_to_system_role( self, messages: list[AllMessageValues], diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py index f75f9159070..ac5032c08f1 100644 --- a/litellm/llms/base_llm/passthrough/transformation.py +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -122,7 +122,7 @@ class RawBytesStreamCollector: self._raw_bytes.append(chunk) def build_logged_response(self, litellm_logging_obj: LiteLLMLoggingObj) -> LoggedRelayResponse | None: - all_chunks: Final = self._provider_config._convert_raw_bytes_to_str_lines(self._raw_bytes) + all_chunks: Final = self._provider_config.convert_raw_bytes_to_str_lines(self._raw_bytes) return self._provider_config.handle_logging_collected_chunks( all_chunks=all_chunks, litellm_logging_obj=litellm_logging_obj, @@ -258,3 +258,9 @@ class BasePassthroughConfig(BaseLLMModelInfo): lines: Final = [line.strip() for line in combined_str.split("\n") if line.strip()] return lines + + def convert_raw_bytes_to_str_lines( + self, + raw_bytes: list[bytes], # mutable-ok: mirrors override contract + ) -> list[str]: # mutable-ok: mirrors override contract + return self._convert_raw_bytes_to_str_lines(raw_bytes) diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index c79b7b6e07d..424e1dc2106 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -790,6 +790,14 @@ class BaseAWSLLM(SignsRequestsWithAWS): self._validate_aws_region_name(aws_region_name) return aws_region_name + def get_aws_region_name( + self, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + model: str | None = None, + model_id: str | None = None, + ) -> str: + return self._get_aws_region_name(optional_params, model, model_id) + @staticmethod def _validate_aws_region_name(aws_region_name: str | None) -> None: """ @@ -804,6 +812,13 @@ class BaseAWSLLM(SignsRequestsWithAWS): "Region names must contain only lowercase letters, digits, and hyphens." ) + @classmethod + def validate_aws_region_name( + cls, + aws_region_name: str | None, + ) -> None: + return cls._validate_aws_region_name(aws_region_name) + @staticmethod def _parse_sts_region_from_endpoint( aws_sts_endpoint: str | None, diff --git a/litellm/llms/bedrock/batches/handler.py b/litellm/llms/bedrock/batches/handler.py index e87c5c35fed..4952b6950af 100644 --- a/litellm/llms/bedrock/batches/handler.py +++ b/litellm/llms/bedrock/batches/handler.py @@ -167,7 +167,7 @@ class BedrockBatchesHandler: ) def job_status() -> "LiteLLMBatch": - return BedrockBatchesHandler._handle_model_invocation_job_status( + return BedrockBatchesHandler.handle_model_invocation_job_status( batch_id=batch_id, aws_region_name=region, logging_obj=logging_obj, @@ -187,7 +187,7 @@ class BedrockBatchesHandler: return job_status() @staticmethod - def _handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj=None, **kwargs) -> "LiteLLMBatch": + def handle_async_invoke_status(batch_id: str, aws_region_name: str, logging_obj=None, **kwargs) -> "LiteLLMBatch": """ Handle async invoke status check for AWS Bedrock. @@ -210,7 +210,7 @@ class BedrockBatchesHandler: embedding_handler: Final = BedrockEmbedding() # Get the status of the async invoke job - status_response: Final = await embedding_handler._get_async_invoke_status( + status_response: Final = await embedding_handler.get_async_invoke_status( invocation_arn=batch_id, aws_region_name=aws_region_name, logging_obj=logging_obj, @@ -263,8 +263,10 @@ class BedrockBatchesHandler: future: Final = executor.submit(run_in_thread) return future.result() + _handle_async_invoke_status = handle_async_invoke_status + @staticmethod - def _handle_model_invocation_job_status( + def handle_model_invocation_job_status( batch_id: str, aws_region_name: str | None = None, logging_obj=None, @@ -406,3 +408,5 @@ class BedrockBatchesHandler: input_file_id=input_uri, output_file_id=output_file_uri if openai_status == "completed" else None, ) + + _handle_model_invocation_job_status = handle_model_invocation_job_status diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index e4001566b8c..3ad8049ec26 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -507,6 +507,13 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): expires_at, ) + def parse_timestamps_and_status( + self, + response_data: Mapping[str, object], + status_str: str, + ) -> tuple[int | None, int | None, int | None, int | None, int | None, int | None]: + return self._parse_timestamps_and_status(response_data, status_str) + def _extract_file_configs(self, response_data): """Helper to extract input and output file configurations.""" # Extract input file ID diff --git a/litellm/llms/bedrock/chat/agentcore/transformation.py b/litellm/llms/bedrock/chat/agentcore/transformation.py index 27102f4289f..207ed249866 100644 --- a/litellm/llms/bedrock/chat/agentcore/transformation.py +++ b/litellm/llms/bedrock/chat/agentcore/transformation.py @@ -193,6 +193,12 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): verbose_logger.debug("Generated new session ID: %s", generated_id) return generated_id + def get_runtime_session_id( + self, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> str: + return self._get_runtime_session_id(optional_params) + def _get_runtime_user_id(self, optional_params: dict) -> str | None: """ Get runtime user ID if provided @@ -202,6 +208,12 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): verbose_logger.debug("Using provided runtimeUserId: %s", user_id) return user_id + def get_runtime_user_id( + self, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> str | None: + return self._get_runtime_user_id(optional_params) + def transform_request( self, model: str, @@ -652,11 +664,11 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): """ from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, - _get_httpx_client, + get_httpx_client, ) if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={}) + client = get_httpx_client(params={}) verbose_logger.debug("Making sync streaming request to: %s", api_base) diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index d54fb1b804f..021327eaa1f 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -12,14 +12,14 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, pop_aws_auth_params, run_aws_signing -from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text +from ..common_utils import BedrockError, error_response_text, get_all_bedrock_regions from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call @@ -37,7 +37,7 @@ def make_sync_call( timeout: float | httpx.Timeout | None = None, ) -> tuple[Any, httpx.Headers]: if client is None: - client = _get_httpx_client() # Create a new client if none provided + client = get_httpx_client() # Create a new client if none provided response: Final = client.post( api_base, @@ -302,7 +302,7 @@ class BedrockConverseLLM(BaseAWSLLM): # and capture it so it can be used as aws_region_name below. _region_from_model: str | None = None _potential_region: Final = _stripped.split("/", 1)[0] - if _potential_region in _get_all_bedrock_regions() and "/" in _stripped: + if _potential_region in get_all_bedrock_regions() and "/" in _stripped: _region_from_model = _potential_region _stripped = _stripped.split("/", 1)[1] _model_for_id = _stripped @@ -443,7 +443,7 @@ class BedrockConverseLLM(BaseAWSLLM): if isinstance(timeout, float) or isinstance(timeout, int): timeout = httpx.Timeout(timeout) _params["timeout"] = timeout - client = _get_httpx_client(_params) + client = get_httpx_client(_params) else: client = client diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 38c47b0c8d7..d78936f6fbc 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -552,7 +552,7 @@ class AmazonConverseConfig(BaseConfig): reasoning_config: Final = self._transform_reasoning_effort_to_reasoning_config(reasoning_effort) optional_params.update(reasoning_config) else: - mapped_thinking: Final = AnthropicConfig._map_reasoning_effort( + mapped_thinking: Final = AnthropicConfig.map_reasoning_effort( reasoning_effort=reasoning_effort, model=model, custom_llm_provider="bedrock", @@ -563,10 +563,10 @@ class AmazonConverseConfig(BaseConfig): optional_params.pop("output_config", None) else: optional_params["thinking"] = mapped_thinking - if AnthropicConfig._is_adaptive_thinking_model(model, "bedrock"): + if AnthropicConfig.is_adaptive_thinking_model(model, "bedrock"): mapped_effort = REASONING_EFFORT_TO_OUTPUT_CONFIG_EFFORT.get(reasoning_effort) if mapped_effort is None: - AnthropicConfig._raise_invalid_reasoning_effort( + AnthropicConfig.raise_invalid_reasoning_effort( model=model, value=reasoning_effort, llm_provider="bedrock_converse", @@ -598,7 +598,7 @@ class AmazonConverseConfig(BaseConfig): model=model, llm_provider="bedrock_converse", ) - error = AnthropicConfig._validate_effort_for_model(model=model, effort=effort, custom_llm_provider="bedrock") + error = AnthropicConfig.validate_effort_for_model(model=model, effort=effort, custom_llm_provider="bedrock") if error is not None: raise litellm.exceptions.BadRequestError( message=error, @@ -1062,7 +1062,7 @@ class AmazonConverseConfig(BaseConfig): optional_params["stopSequences"] = value if param == "temperature" or param == "top_p": if base_model.startswith("anthropic"): - AnthropicConfig._apply_sampling_param( + AnthropicConfig.apply_sampling_param( optional_params=optional_params, model=model, param=param, @@ -1109,10 +1109,10 @@ class AmazonConverseConfig(BaseConfig): if ( isinstance(value, dict) and value.get("type") == "adaptive" - and not AnthropicConfig._is_adaptive_thinking_model(model, "bedrock") + and not AnthropicConfig.is_adaptive_thinking_model(model, "bedrock") ): max_tokens = non_default_params.get("max_completion_tokens") or non_default_params.get("max_tokens") - legacy_thinking = AnthropicConfig._map_reasoning_effort( + legacy_thinking = AnthropicConfig.map_reasoning_effort( reasoning_effort="medium", model=model, custom_llm_provider="bedrock", @@ -1557,7 +1557,7 @@ class AmazonConverseConfig(BaseConfig): if val_top_k is not None: if base_model.startswith("anthropic"): top_k_params: Final[dict] = {} - AnthropicConfig._apply_sampling_param( + AnthropicConfig.apply_sampling_param( optional_params=top_k_params, model=model, param="top_k", @@ -1695,7 +1695,7 @@ class AmazonConverseConfig(BaseConfig): if is_bedrock_application_inference_profile_arn(model): additional_request_params["output_config"] = anthropic_output_config elif base_model.startswith("anthropic"): - if litellm.drop_params is True and not AnthropicConfig._model_supports_effort_param(model, "bedrock"): + if litellm.drop_params is True and not AnthropicConfig.model_supports_effort_param(model, "bedrock"): litellm.verbose_logger.warning( DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING, model, @@ -1846,7 +1846,7 @@ class AmazonConverseConfig(BaseConfig): if ( isinstance(output_config, dict) and output_config.get("effort") is not None - and not AnthropicConfig._is_adaptive_thinking_model(model, "bedrock") + and not AnthropicConfig.is_adaptive_thinking_model(model, "bedrock") ): from litellm.types.llms.anthropic import ( ANTHROPIC_EFFORT_BETA_HEADER, diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 8a2291d33be..6276324ad7e 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -18,8 +18,8 @@ from litellm.llms.anthropic.chat.handler import ( from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.bedrock import * from litellm.types.llms.openai import ( @@ -297,7 +297,7 @@ def make_sync_call( ) -> "tuple[MockResponseIterator | Iterator[GChunk | ModelResponseStream | dict], httpx.Headers]": try: if client is None: - client = _get_httpx_client( + client = get_httpx_client( params=( {"ssl_verify": logging_obj.litellm_params.get("ssl_verify")} if logging_obj and logging_obj.litellm_params and logging_obj.litellm_params.get("ssl_verify") @@ -777,6 +777,12 @@ class AWSEventStreamDecoder: tool_use=None, ) + def chunk_parser( + self, + chunk_data: dict[str, object], # mutable-ok: mirrors override contract + ) -> GChunk | ModelResponseStream | dict[str, object]: # mutable-ok: mirrors override contract + return self._chunk_parser(chunk_data) + def iter_bytes( self, iterator: Iterator[bytes], *, response_headers: Mapping[str, str] | None = None ) -> Iterator[GChunk | ModelResponseStream | dict]: @@ -793,7 +799,7 @@ class AWSEventStreamDecoder: if message: # sse_event = ServerSentEvent(data=message, event="completion") _data = json.loads(message) - yield self._chunk_parser(chunk_data=_data) + yield self.chunk_parser(chunk_data=_data) undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers) if undecoded_stream_error is not None: raise undecoded_stream_error @@ -813,7 +819,7 @@ class AWSEventStreamDecoder: message = self._decode_event(event, tally) if message: _data = json.loads(message) - yield self._chunk_parser(chunk_data=_data) + yield self.chunk_parser(chunk_data=_data) undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers) if undecoded_stream_error is not None: raise undecoded_stream_error @@ -955,7 +961,7 @@ class MockResponseIterator: # for returning ai21 streaming responses """ tool_use: ChatCompletionToolCallChunk | None = None if self.json_mode is True and tool_calls is not None: - message: Final = litellm.AnthropicConfig()._convert_tool_response_to_message(tool_calls=tool_calls) + message: Final = litellm.AnthropicConfig().convert_tool_response_to_message(tool_calls=tool_calls) if message is not None: text = message.content or "" tool_use = None diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index a1fa7963ab5..af833d93084 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -28,7 +28,7 @@ from litellm.llms.bedrock.request_metadata import ( from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, + get_httpx_client, ) from litellm.types.llms.bedrock import GuardrailConfigBlock from litellm.types.llms.openai import AllMessageValues @@ -499,7 +499,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: sync_client: Final = ( - _get_httpx_client(params={}) if client is None or isinstance(client, AsyncHTTPHandler) else client + get_httpx_client(params={}) if client is None or isinstance(client, AsyncHTTPHandler) else client ) chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size completion_stream, response_headers = make_sync_call( diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 0352234ecb8..7c637bf6724 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -313,7 +313,7 @@ def strip_unsupported_bedrock_invoke_output_config_keys( return if all(key == "format" for key in output_config): return - if _bedrock_model_supports(model, "supports_output_config") or AnthropicConfig._model_supports_effort_param( + if _bedrock_model_supports(model, "supports_output_config") or AnthropicConfig.model_supports_effort_param( model, "bedrock" ): return @@ -773,7 +773,7 @@ def get_bedrock_tool_name(response_tool_name: str) -> str: _BEDROCK_GLOBAL_REGIONS: list[str] | None = None -def _get_all_bedrock_regions() -> list[str]: +def get_all_bedrock_regions() -> list[str]: """Get all Bedrock regions, cached at module level.""" global _BEDROCK_GLOBAL_REGIONS if _BEDROCK_GLOBAL_REGIONS is None: @@ -781,6 +781,9 @@ def _get_all_bedrock_regions() -> list[str]: return _BEDROCK_GLOBAL_REGIONS +_get_all_bedrock_regions = get_all_bedrock_regions + + def get_bedrock_cross_region_inference_regions() -> list[str]: """Abbreviations of regions AWS Bedrock supports for cross region inference.""" return ["global", "us", "eu", "apac", "jp", "au", "us-gov"] @@ -829,7 +832,7 @@ def split_bedrock_region_path(model: str) -> tuple[str | None, str]: """ stripped: Final = strip_bedrock_routing_prefix(model) region, separator, model_id = stripped.partition("/") - if separator and region in _get_all_bedrock_regions(): + if separator and region in get_all_bedrock_regions(): return region, model_id return None, stripped @@ -1104,7 +1107,7 @@ def get_bedrock_base_model(model: str) -> str: if potential_region in get_bedrock_cross_region_inference_regions(): return model.split(".", 1)[1] - elif alt_potential_region in _get_all_bedrock_regions() and len(model.split("/", 1)) > 1: + elif alt_potential_region in get_all_bedrock_regions() and len(model.split("/", 1)) > 1: return model.split("/", 1)[1] return model @@ -1174,7 +1177,7 @@ def bedrock_supports_tool_search(model: str) -> bool: """ from litellm.llms.anthropic.common_utils import AnthropicModelInfo - return AnthropicModelInfo._supports_model_capability(model, "supports_tool_search", "bedrock") + return AnthropicModelInfo.supports_model_capability(model, "supports_tool_search", "bedrock") def is_claude_4_5_on_bedrock(model: str) -> bool: @@ -2055,7 +2058,7 @@ class CommonBatchFilesUtils: except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") - aws_region_name: Final = self._base_aws._get_aws_region_name(optional_params=optional_params, model="") + aws_region_name: Final = self._base_aws.get_aws_region_name(optional_params=optional_params, model="") credentials: Final = self._base_aws.resolve_credentials( AwsAuthParams.model_validate(optional_params), aws_region_name ) diff --git a/litellm/llms/bedrock/embed/amazon_nova_transformation.py b/litellm/llms/bedrock/embed/amazon_nova_transformation.py index 2709878e489..cc8c8d37c42 100644 --- a/litellm/llms/bedrock/embed/amazon_nova_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_nova_transformation.py @@ -204,6 +204,16 @@ class AmazonNovaEmbeddingConfig: return request + def transform_request( + self, + input: str, + inference_params: dict[str, object], # mutable-ok: mirrors override contract + async_invoke_route: bool = False, + model_id: str | None = None, + output_s3_uri: str | None = None, + ) -> dict[str, object]: # mutable-ok: mirrors override contract + return self._transform_request(input, inference_params, async_invoke_route, model_id, output_s3_uri) + def _wrap_async_invoke_request( self, model_input: dict, @@ -316,6 +326,14 @@ class AmazonNovaEmbeddingConfig: return EmbeddingResponse(data=embeddings, model=model, usage=usage) + def transform_response( + self, + response_list: list[dict[str, object]], # mutable-ok: mirrors override contract + model: str, + batch_data: list[dict[str, object]] | None = None, # mutable-ok: mirrors override contract + ) -> EmbeddingResponse: + return self._transform_response(response_list, model, batch_data) + def _transform_async_invoke_response(self, response: dict, model: str) -> EmbeddingResponse: """ Transform async invoke response (invocation ARN) to OpenAI format. @@ -351,3 +369,10 @@ class AmazonNovaEmbeddingConfig: usage=usage, hidden_params=hidden_params, ) + + def transform_async_invoke_response( + self, + response: dict[str, object], # mutable-ok: mirrors override contract + model: str, + ) -> EmbeddingResponse: + return self._transform_async_invoke_response(response, model) diff --git a/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py index 143e0b2f623..581b89369f7 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py @@ -60,6 +60,13 @@ class AmazonTitanG1Config: def _transform_request(self, input: str, inference_params: dict) -> AmazonTitanG1EmbeddingRequest: return AmazonTitanG1EmbeddingRequest(inputText=input) + def transform_request( + self, + input: str, + inference_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> AmazonTitanG1EmbeddingRequest: + return self._transform_request(input, inference_params) + def _transform_response(self, response_list: list[dict], model: str) -> EmbeddingResponse: total_prompt_tokens = 0 @@ -81,3 +88,10 @@ class AmazonTitanG1Config: total_tokens=total_prompt_tokens, ) return EmbeddingResponse(model=model, usage=usage, data=transformed_responses) + + def transform_response( + self, + response_list: list[dict[str, object]], # mutable-ok: mirrors override contract + model: str, + ) -> EmbeddingResponse: + return self._transform_response(response_list, model) diff --git a/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py index 5897ad84115..e7a5b4a3d78 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py @@ -52,6 +52,13 @@ class AmazonTitanMultimodalEmbeddingG1Config: transformed_request[k] = v return transformed_request + def transform_request( + self, + input: str, + inference_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> AmazonTitanMultimodalEmbeddingRequest: + return self._transform_request(input, inference_params) + def _transform_response( self, response_list: list[dict], @@ -91,3 +98,11 @@ class AmazonTitanMultimodalEmbeddingG1Config: prompt_tokens_details=prompt_tokens_details, ) return EmbeddingResponse(model=model, usage=usage, data=transformed_responses) + + def transform_response( + self, + response_list: list[dict[str, object]], # mutable-ok: mirrors override contract + model: str, + batch_data: list[dict[str, object]] | None = None, # mutable-ok: mirrors override contract + ) -> EmbeddingResponse: + return self._transform_response(response_list, model, batch_data) diff --git a/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py b/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py index 3f70bbc3bc3..3b594662a9b 100644 --- a/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py +++ b/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py @@ -77,6 +77,13 @@ class AmazonTitanV2Config: def _transform_request(self, input: str, inference_params: dict) -> AmazonTitanV2EmbeddingRequest: return AmazonTitanV2EmbeddingRequest(inputText=input, **inference_params) + def transform_request( + self, + input: str, + inference_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> AmazonTitanV2EmbeddingRequest: + return self._transform_request(input, inference_params) + def _transform_response(self, response_list: list[dict], model: str) -> EmbeddingResponse: total_prompt_tokens = 0 @@ -116,3 +123,10 @@ class AmazonTitanV2Config: total_tokens=total_prompt_tokens, ) return EmbeddingResponse(model=model, usage=usage, data=transformed_responses) + + def transform_response( + self, + response_list: list[dict[str, object]], # mutable-ok: mirrors override contract + model: str, + ) -> EmbeddingResponse: + return self._transform_response(response_list, model) diff --git a/litellm/llms/bedrock/embed/cohere_transformation.py b/litellm/llms/bedrock/embed/cohere_transformation.py index 8a17bb9d595..9b626515271 100644 --- a/litellm/llms/bedrock/embed/cohere_transformation.py +++ b/litellm/llms/bedrock/embed/cohere_transformation.py @@ -41,3 +41,11 @@ class BedrockCohereEmbeddingConfig: new_transformed_request[k] = transformed_request[k] return new_transformed_request + + def transform_request( + self, + model: str, + input: list[str], # mutable-ok: mirrors override contract + inference_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> CohereEmbeddingRequest: + return self._transform_request(model, input, inference_params) diff --git a/litellm/llms/bedrock/embed/embedding.py b/litellm/llms/bedrock/embed/embedding.py index 46d7b1ef9e7..58cfb6cb95e 100644 --- a/litellm/llms/bedrock/embed/embedding.py +++ b/litellm/llms/bedrock/embed/embedding.py @@ -16,13 +16,14 @@ from litellm.llms.cohere.embed.handler import embedding as cohere_embedding from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.secret_managers.main import get_secret from litellm.types.llms.bedrock import ( AmazonEmbeddingRequest, CohereEmbeddingRequest, + TwelveLabsAsyncInvokeStatusResponse, ) from litellm.types.utils import EmbeddingResponse, LlmProviders @@ -122,7 +123,7 @@ class BedrockEmbedding(BaseAWSLLM): if isinstance(timeout, float) or isinstance(timeout, int): timeout = httpx.Timeout(timeout) _params["timeout"] = timeout - client = _get_httpx_client(_params) + client = get_httpx_client(_params) else: client = client try: @@ -191,11 +192,11 @@ class BedrockEmbedding(BaseAWSLLM): # Handle async invoke responses (single response with invocationArn) if is_async_invoke and len(response_list) == 1 and "invocationArn" in response_list[0]: if provider == "twelvelabs": - returned_response = TwelveLabsMarengoEmbeddingConfig()._transform_async_invoke_response( + returned_response = TwelveLabsMarengoEmbeddingConfig().transform_async_invoke_response( response=response_list[0], model=model ) elif provider == "nova": - returned_response = AmazonNovaEmbeddingConfig()._transform_async_invoke_response( + returned_response = AmazonNovaEmbeddingConfig().transform_async_invoke_response( response=response_list[0], model=model ) else: @@ -226,21 +227,21 @@ class BedrockEmbedding(BaseAWSLLM): else: # Handle regular invoke responses if model == "amazon.titan-embed-image-v1": - returned_response = AmazonTitanMultimodalEmbeddingG1Config()._transform_response( + returned_response = AmazonTitanMultimodalEmbeddingG1Config().transform_response( response_list=response_list, model=model, batch_data=batch_data ) elif model == "amazon.titan-embed-text-v1": - returned_response = AmazonTitanG1Config()._transform_response(response_list=response_list, model=model) + returned_response = AmazonTitanG1Config().transform_response(response_list=response_list, model=model) elif model == "amazon.titan-embed-text-v2:0": - returned_response = AmazonTitanV2Config()._transform_response(response_list=response_list, model=model) + returned_response = AmazonTitanV2Config().transform_response(response_list=response_list, model=model) elif model == "amazon.titan-embed-g1-text-02": - returned_response = AmazonTitanG1Config()._transform_response(response_list=response_list, model=model) + returned_response = AmazonTitanG1Config().transform_response(response_list=response_list, model=model) elif provider == "twelvelabs": - returned_response = TwelveLabsMarengoEmbeddingConfig()._transform_response( + returned_response = TwelveLabsMarengoEmbeddingConfig().transform_response( response_list=response_list, model=model, batch_data=batch_data ) elif provider == "nova": - returned_response = AmazonNovaEmbeddingConfig()._transform_response( + returned_response = AmazonNovaEmbeddingConfig().transform_response( response_list=response_list, model=model, batch_data=batch_data ) @@ -438,7 +439,7 @@ class BedrockEmbedding(BaseAWSLLM): data: CohereEmbeddingRequest | None = None batch_data: list | None = None if provider == "cohere": - data = BedrockCohereEmbeddingConfig()._transform_request( + data = BedrockCohereEmbeddingConfig().transform_request( model=model, input=input, inference_params=inference_params ) elif provider == "amazon" and model in [ @@ -451,20 +452,20 @@ class BedrockEmbedding(BaseAWSLLM): for i in input: if model == "amazon.titan-embed-image-v1": transformed_request: AmazonEmbeddingRequest = ( - AmazonTitanMultimodalEmbeddingG1Config()._transform_request( + AmazonTitanMultimodalEmbeddingG1Config().transform_request( input=i, inference_params=inference_params ) ) elif model == "amazon.titan-embed-text-v1": - transformed_request = AmazonTitanG1Config()._transform_request( + transformed_request = AmazonTitanG1Config().transform_request( input=i, inference_params=inference_params ) elif model == "amazon.titan-embed-text-v2:0": - transformed_request = AmazonTitanV2Config()._transform_request( + transformed_request = AmazonTitanV2Config().transform_request( input=i, inference_params=inference_params ) elif model == "amazon.titan-embed-g1-text-02": - transformed_request = AmazonTitanG1Config()._transform_request( + transformed_request = AmazonTitanG1Config().transform_request( input=i, inference_params=inference_params ) else: @@ -483,7 +484,7 @@ class BedrockEmbedding(BaseAWSLLM): elif provider == "twelvelabs": batch_data = [] for i in input: - twelvelabs_request = TwelveLabsMarengoEmbeddingConfig(model=model)._transform_request( + twelvelabs_request = TwelveLabsMarengoEmbeddingConfig(model=model).transform_request( input=i, inference_params=inference_params, async_invoke_route=has_async_invoke, @@ -495,7 +496,7 @@ class BedrockEmbedding(BaseAWSLLM): elif provider == "nova": batch_data = [] for i in input: - nova_request = AmazonNovaEmbeddingConfig()._transform_request( + nova_request = AmazonNovaEmbeddingConfig().transform_request( input=i, inference_params=inference_params, async_invoke_route=has_async_invoke, @@ -665,3 +666,12 @@ class BedrockEmbedding(BaseAWSLLM): return response.json() else: raise Exception(f"Failed to get async invoke status: {response.status_code} - {response.text}") + + async def get_async_invoke_status( + self, + invocation_arn: str, + aws_region_name: str, + logging_obj: "LiteLLMLoggingObj | None" = None, + **kwargs: object, # kwargs-ok: mirrors private method extension kwargs + ) -> TwelveLabsAsyncInvokeStatusResponse: + return await self._get_async_invoke_status(invocation_arn, aws_region_name, logging_obj, **kwargs) diff --git a/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py b/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py index 9391068b8aa..49a35e7cdc0 100644 --- a/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py +++ b/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py @@ -272,6 +272,19 @@ class TwelveLabsMarengoEmbeddingConfig: return transformed_request + def transform_request( + self, + input: str, + inference_params: dict[str, object], # mutable-ok: mirrors override contract + async_invoke_route: bool = False, + model_id: str | None = None, + output_s3_uri: str | None = None, + drop_params: bool = False, + ) -> TwelveLabsMarengoEmbeddingRequest | TwelveLabsMarengo3EmbeddingRequest | TwelveLabsAsyncInvokeRequest: + return self._transform_request( + input, inference_params, async_invoke_route, model_id, output_s3_uri, drop_params + ) + def _wrap_async_invoke_request( self, model_input: TwelveLabsMarengoEmbeddingRequest | TwelveLabsMarengo3EmbeddingRequest, @@ -319,6 +332,14 @@ class TwelveLabsMarengoEmbeddingConfig: ] return EmbeddingResponse(data=embeddings, model=model, usage=_billed_usage(batch_data)) + def transform_response( + self, + response_list: list[dict[str, object]], # mutable-ok: mirrors override contract + model: str, + batch_data: list[dict[str, object]] | None = None, # mutable-ok: mirrors override contract + ) -> EmbeddingResponse: + return self._transform_response(response_list, model, batch_data) + def _transform_async_invoke_response(self, response: dict, model: str) -> EmbeddingResponse: """ Transform async invoke response (invocation ARN) to OpenAI format. @@ -366,3 +387,10 @@ class TwelveLabsMarengoEmbeddingConfig: usage=usage, hidden_params=hidden_params, ) + + def transform_async_invoke_response( + self, + response: dict[str, object], # mutable-ok: mirrors override contract + model: str, + ) -> EmbeddingResponse: + return self._transform_async_invoke_response(response, model) diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index 04969027638..e7ee2f89bbf 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -817,7 +817,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): non_default_params=non_default_params, optional_params={}, ) - return dict(titan_config._transform_request(input=input_text, inference_params=inference_params)) + return dict(titan_config.transform_request(input=input_text, inference_params=inference_params)) @staticmethod def _transform_text_completion_body_to_chat_body( @@ -1024,6 +1024,13 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): bedrock_jsonl_content.append(bedrock_record) return bedrock_jsonl_content + def transform_openai_jsonl_content_to_bedrock_jsonl_content( + self, + openai_jsonl_content: Sequence[_OpenAIBatchRecord], + target_model: str = "", + ) -> list[_BedrockBatchRecord]: # mutable-ok: mirrors override contract + return self._transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content, target_model) + def transform_create_file_request( self, model: str, @@ -1535,7 +1542,7 @@ class BedrockJsonlFilesTransformation: Delegate to the main BedrockFilesConfig transformation method """ config: Final = BedrockFilesConfig() - return config._transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content) + return config.transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content) def _get_s3_object_name( self, diff --git a/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py b/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py index f94f4a3657c..db0dc7e65c2 100644 --- a/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py +++ b/litellm/llms/bedrock/image_edit/amazon_nova_canvas_image_edit_transformation.py @@ -229,6 +229,13 @@ class BedrockAmazonNovaCanvasImageEditConfig(BaseImageEditConfig): """ return _supports_nova_canvas_image_edit_from_model_cost(model or "") + @classmethod + def is_nova_canvas_image_edit_model( + cls, + model: str | None = None, + ) -> bool: + return cls._is_nova_canvas_image_edit_model(model) + def get_error_class( self, error_message: str, @@ -509,9 +516,9 @@ def get_bedrock_image_edit_config_for_model( BedrockStabilityImageEditConfig, ) - if BedrockStabilityImageEditConfig._is_stability_edit_model(model): + if BedrockStabilityImageEditConfig.is_stability_edit_model(model): return BedrockStabilityImageEditConfig() - if BedrockAmazonNovaCanvasImageEditConfig._is_nova_canvas_image_edit_model(model): + if BedrockAmazonNovaCanvasImageEditConfig.is_nova_canvas_image_edit_model(model): return BedrockAmazonNovaCanvasImageEditConfig() raise ValueError( f"Unsupported Bedrock image-edit model: {model!r}. " diff --git a/litellm/llms/bedrock/image_edit/handler.py b/litellm/llms/bedrock/image_edit/handler.py index 58de67107f4..76ec537b68d 100644 --- a/litellm/llms/bedrock/image_edit/handler.py +++ b/litellm/llms/bedrock/image_edit/handler.py @@ -23,8 +23,8 @@ from litellm.llms.bedrock.image_edit.stability_transformation import ( from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.utils import ImageResponse @@ -56,9 +56,9 @@ class BedrockImageEdit(BaseAWSLLM): @classmethod def get_config_class(cls, model: str | None): - if BedrockStabilityImageEditConfig._is_stability_edit_model(model): + if BedrockStabilityImageEditConfig.is_stability_edit_model(model): return BedrockStabilityImageEditConfig - if BedrockAmazonNovaCanvasImageEditConfig._is_nova_canvas_image_edit_model(model): + if BedrockAmazonNovaCanvasImageEditConfig.is_nova_canvas_image_edit_model(model): return BedrockAmazonNovaCanvasImageEditConfig raise ValueError( f"Unsupported Bedrock image-edit model: {model!r}. " @@ -104,7 +104,7 @@ class BedrockImageEdit(BaseAWSLLM): ) if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client() + client = get_httpx_client() try: response: Final = client.post( url=prepared_request.endpoint_url, diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index dc8932b76d0..d0ce5895970 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -86,6 +86,13 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): return True return False + @classmethod + def is_stability_edit_model( + cls, + model: str | None = None, + ) -> bool: + return cls._is_stability_edit_model(model) + def get_error_class( self, error_message: str, diff --git a/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py b/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py index ce61a6253f6..63725473923 100644 --- a/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_nova_canvas_transformation.py @@ -59,6 +59,13 @@ class AmazonNovaCanvasConfig: return True return False + @classmethod + def is_nova_model( + cls, + model: str | None = None, + ) -> bool: + return cls._is_nova_model(model) + @classmethod def transform_request_body(cls, text: str, optional_params: dict) -> AmazonNovaCanvasRequestBase: """ diff --git a/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py b/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py index 98e4cbbfd4d..acc847e0e3c 100644 --- a/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_stability3_transformation.py @@ -65,6 +65,13 @@ class AmazonStability3Config: return True return False + @classmethod + def is_stability_3_model( + cls, + model: str | None = None, + ) -> bool: + return cls._is_stability_3_model(model) + @classmethod def transform_request_body(cls, text: str, optional_params: dict) -> AmazonStability3TextToImageRequest: """ diff --git a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py index 12ace32f43b..145e6a0781f 100644 --- a/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py +++ b/litellm/llms/bedrock/image_generation/amazon_titan_transformation.py @@ -70,6 +70,13 @@ class AmazonTitanImageGenerationConfig: return True return False + @classmethod + def is_titan_model( + cls, + model: str | None = None, + ) -> bool: + return cls._is_titan_model(model) + @classmethod def get_supported_openai_params(cls, model: str | None = None) -> list: return ["size", "n", "quality"] diff --git a/litellm/llms/bedrock/image_generation/image_handler.py b/litellm/llms/bedrock/image_generation/image_handler.py index bdf9f596053..018e92fd79f 100644 --- a/litellm/llms/bedrock/image_generation/image_handler.py +++ b/litellm/llms/bedrock/image_generation/image_handler.py @@ -23,8 +23,8 @@ from litellm.llms.bedrock.image_generation.amazon_titan_transformation import ( from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.base import LiteLLMBaseModel from litellm.types.utils import ImageResponse @@ -64,11 +64,11 @@ class BedrockImageGeneration(BaseAWSLLM): @classmethod def get_config_class(cls, model: str | None) -> BedrockImageConfigClass: - if AmazonTitanImageGenerationConfig._is_titan_model(model): + if AmazonTitanImageGenerationConfig.is_titan_model(model): return AmazonTitanImageGenerationConfig - elif AmazonNovaCanvasConfig._is_nova_model(model): + elif AmazonNovaCanvasConfig.is_nova_model(model): return AmazonNovaCanvasConfig - elif AmazonStability3Config._is_stability_3_model(model): + elif AmazonStability3Config.is_stability_3_model(model): return AmazonStability3Config else: return litellm.AmazonStabilityConfig @@ -109,7 +109,7 @@ class BedrockImageGeneration(BaseAWSLLM): ) if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client() + client = get_httpx_client() try: response: Final = client.post( url=prepared_request.endpoint_url, diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 02f02d97e63..26a0302ef61 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -226,7 +226,7 @@ class AmazonAnthropicClaudeMessagesConfig( Returns: True if the model supports extended thinking on Bedrock """ - if AnthropicModelInfo._is_adaptive_thinking_model(model, "bedrock"): + if AnthropicModelInfo.is_adaptive_thinking_model(model, "bedrock"): return True model_lower: Final = model.lower() @@ -276,7 +276,7 @@ class AmazonAnthropicClaudeMessagesConfig( if not self._supports_extended_thinking_on_bedrock(model): return False - is_adaptive_thinking_model: Final = AnthropicModelInfo._is_adaptive_thinking_model(model, "bedrock") + is_adaptive_thinking_model: Final = AnthropicModelInfo.is_adaptive_thinking_model(model, "bedrock") thinking: Final = anthropic_messages_request.get("thinking") if isinstance(thinking, dict): @@ -664,7 +664,7 @@ class AmazonAnthropicClaudeMessagesConfig( path degrades ``xhigh`` -> ``max`` rather than 400-ing. Non-adaptive models and models without a ceiling are left untouched. """ - if not AnthropicModelInfo._is_adaptive_thinking_model(model, "bedrock"): + if not AnthropicModelInfo.is_adaptive_thinking_model(model, "bedrock"): return effort: Final = optional_params.get("reasoning_effort") if not isinstance(effort, str): @@ -778,7 +778,7 @@ class AmazonAnthropicClaudeMessagesConfig( litellm.drop_params is True and isinstance(remaining_output_config, dict) and any(key != "format" for key in remaining_output_config) - and not AnthropicConfig._model_supports_effort_param(model, "bedrock") + and not AnthropicConfig.model_supports_effort_param(model, "bedrock") ): verbose_logger.warning( DROP_UNSUPPORTED_OUTPUT_CONFIG_WARNING, diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index 6a120f41cb6..c61576fa7b2 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -74,7 +74,7 @@ def _translate_message(decoder: "AWSEventStreamDecoder", message: str) -> ModelR ) from litellm.types.utils import GenericStreamingChunk - translated_chunk: Final = decoder._chunk_parser(chunk_data=json.loads(message)) + translated_chunk: Final = decoder.chunk_parser(chunk_data=json.loads(message)) if isinstance(translated_chunk, ModelResponseStream): return translated_chunk if generic_chunk_has_all_required_fields(cast(dict, translated_chunk)): diff --git a/litellm/llms/bedrock/rerank/handler.py b/litellm/llms/bedrock/rerank/handler.py index 4a8308bb46a..01ef5af2b67 100644 --- a/litellm/llms/bedrock/rerank/handler.py +++ b/litellm/llms/bedrock/rerank/handler.py @@ -9,8 +9,8 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.bedrock import BedrockPreparedRequest from litellm.types.rerank import RerankRequest @@ -58,7 +58,7 @@ class BedrockRerankHandler(BaseAWSLLM): except httpx.TimeoutException: raise BedrockError(status_code=408, message="Timeout error occurred.") - return BedrockRerankConfig()._transform_response(_JSON_DICT.validate_python(response.json())) + return BedrockRerankConfig().transform_response(_JSON_DICT.validate_python(response.json())) def rerank( self, @@ -85,7 +85,7 @@ class BedrockRerankHandler(BaseAWSLLM): rank_fields=rank_fields, return_documents=return_documents, ) - data: Final = BedrockRerankConfig()._transform_request(request_data) + data: Final = BedrockRerankConfig().transform_request(request_data) prepared_request: Final = self._prepare_request( model=model, @@ -114,7 +114,7 @@ class BedrockRerankHandler(BaseAWSLLM): ) if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client() + client = get_httpx_client() try: response: Final = client.post( url=prepared_request["endpoint_url"], @@ -141,7 +141,7 @@ class BedrockRerankHandler(BaseAWSLLM): response_json: Final = _JSON_DICT.validate_python(response.json()) - return BedrockRerankConfig()._transform_response(response_json) + return BedrockRerankConfig().transform_response(response_json) def _prepare_request( self, diff --git a/litellm/llms/bedrock/rerank/transformation.py b/litellm/llms/bedrock/rerank/transformation.py index 930dfa4dfbd..442903fc05e 100644 --- a/litellm/llms/bedrock/rerank/transformation.py +++ b/litellm/llms/bedrock/rerank/transformation.py @@ -77,6 +77,12 @@ class BedrockRerankConfig: sources=_sources, ) + def transform_request( + self, + request_data: RerankRequest, + ) -> BedrockRerankRequest: + return self._transform_request(request_data) + def _transform_response(self, response: dict) -> RerankResponse: """ Transform the response from Bedrock into the RerankResponse format. @@ -108,3 +114,9 @@ class BedrockRerankConfig: results=_results, meta=rerank_meta, ) # Return response + + def transform_response( + self, + response: dict[str, object], # mutable-ok: mirrors override contract + ) -> RerankResponse: + return self._transform_response(response) diff --git a/litellm/llms/bedrock_mantle/chat/transformation.py b/litellm/llms/bedrock_mantle/chat/transformation.py index 41d93a8dd4d..65deac0c4c8 100644 --- a/litellm/llms/bedrock_mantle/chat/transformation.py +++ b/litellm/llms/bedrock_mantle/chat/transformation.py @@ -71,7 +71,7 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig): or get_secret_str("AWS_REGION") or BEDROCK_MANTLE_DEFAULT_REGION ) - BaseAWSLLM._validate_aws_region_name(region) + BaseAWSLLM.validate_aws_region_name(region) # The base path segment is data-driven per model (use_openai_responses_path # flag): gemma-4-* and gpt-5.x are served on /openai/v1, everything else on # /v1. An explicit api_base still wins over the derived default. @@ -83,6 +83,15 @@ class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig): dynamic_api_key: Final = self._resolve_bearer_token(api_key) return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + litellm_params: GenericLiteLLMParams | None = None, + model: str | None = None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key, litellm_params, model) + def validate_environment( self, headers: dict, diff --git a/litellm/llms/bedrock_mantle/common_utils.py b/litellm/llms/bedrock_mantle/common_utils.py index 232d6dcf70f..8f9cf5c4f6b 100644 --- a/litellm/llms/bedrock_mantle/common_utils.py +++ b/litellm/llms/bedrock_mantle/common_utils.py @@ -48,7 +48,7 @@ def split_mantle_region_prefix(model: str) -> tuple[str | None, str]: def resolve_mantle_region(params: Mapping[str, object]) -> str: region: Final = params.get("aws_region_name") if isinstance(region, str) and region: - BaseAWSLLM._validate_aws_region_name(region) + BaseAWSLLM.validate_aws_region_name(region) return region api_base: Final = params.get("api_base") base: Final = (api_base if isinstance(api_base, str) else None) or get_secret_str("BEDROCK_MANTLE_API_BASE") diff --git a/litellm/llms/black_forest_labs/image_edit/handler.py b/litellm/llms/black_forest_labs/image_edit/handler.py index 178acb0de0d..e7d0c9ed8ae 100644 --- a/litellm/llms/black_forest_labs/image_edit/handler.py +++ b/litellm/llms/black_forest_labs/image_edit/handler.py @@ -20,8 +20,8 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import FileTypes, ImageResponse @@ -138,7 +138,7 @@ class BlackForestLabsImageEdit: # Sync version if client is None or not isinstance(client, HTTPHandler): - sync_client = _get_httpx_client() + sync_client = get_httpx_client() else: sync_client = client diff --git a/litellm/llms/black_forest_labs/image_generation/handler.py b/litellm/llms/black_forest_labs/image_generation/handler.py index 879bef37b58..043ddc077db 100644 --- a/litellm/llms/black_forest_labs/image_generation/handler.py +++ b/litellm/llms/black_forest_labs/image_generation/handler.py @@ -20,8 +20,8 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ImageResponse @@ -119,7 +119,7 @@ class BlackForestLabsImageGeneration: # Sync version if client is None or not isinstance(client, HTTPHandler): - sync_client = _get_httpx_client() + sync_client = get_httpx_client() else: sync_client = client diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py index 6a5f97e46b4..e107674b462 100644 --- a/litellm/llms/bytez/chat/transformation.py +++ b/litellm/llms/bytez/chat/transformation.py @@ -13,8 +13,8 @@ from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMExcepti from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, version, ) from litellm.types.llms.openai import AllMessageValues @@ -264,7 +264,7 @@ class BytezChatConfig(BaseConfig): timeout: float | httpx.Timeout | None = None, ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): - client = _get_httpx_client(params={}) + client = get_httpx_client(params={}) try: response: Final = client.post( diff --git a/litellm/llms/chatgpt/authenticator.py b/litellm/llms/chatgpt/authenticator.py index 5ab16cd189a..ff9c6476cba 100644 --- a/litellm/llms/chatgpt/authenticator.py +++ b/litellm/llms/chatgpt/authenticator.py @@ -12,7 +12,7 @@ from litellm._logging import verbose_logger from litellm.constants import HTTP_HANDLER_CONNECT_TIMEOUT_SECONDS from litellm.litellm_core_utils.asyncify import can_block_current_thread from litellm.litellm_core_utils.request_timeout_resolver import get_configured_request_timeout -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client from .common_utils import ( CHATGPT_API_BASE, @@ -198,7 +198,7 @@ class Authenticator: def _request_device_code(self) -> dict[str, str]: try: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() resp: Final = client.post( CHATGPT_DEVICE_CODE_URL, json={"client_id": CHATGPT_CLIENT_ID}, @@ -231,7 +231,7 @@ class Authenticator: } def _poll_for_authorization_code(self, device_code: dict[str, str]) -> dict[str, str]: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() interval: Final = int(device_code.get("interval", "5")) start_time: Final = time.time() while time.time() - start_time < DEVICE_CODE_TIMEOUT_SECONDS: @@ -281,7 +281,7 @@ class Authenticator: def _exchange_code_for_tokens(self, code_data: dict[str, str]) -> dict[str, str]: try: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() redirect_uri: Final = f"{CHATGPT_AUTH_BASE}/deviceauth/callback" body: Final = ( "grant_type=authorization_code" @@ -324,7 +324,7 @@ class Authenticator: def _refresh_tokens(self, refresh_token: str) -> dict[str, str]: try: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() resp: Final = client.post( CHATGPT_OAUTH_TOKEN_URL, json={ diff --git a/litellm/llms/chatgpt/chat/transformation.py b/litellm/llms/chatgpt/chat/transformation.py index 1b110704c8b..7c985b7dd5e 100644 --- a/litellm/llms/chatgpt/chat/transformation.py +++ b/litellm/llms/chatgpt/chat/transformation.py @@ -44,6 +44,15 @@ class ChatGPTConfig(OpenAIConfig): ) return dynamic_api_base, dynamic_api_key, custom_llm_provider + def get_openai_compatible_provider_info( + self, + model: str, + api_base: str | None, + api_key: str | None, + custom_llm_provider: str, + ) -> tuple[str | None, str | None, str]: + return self._get_openai_compatible_provider_info(model, api_base, api_key, custom_llm_provider) + def validate_environment( self, headers: dict, diff --git a/litellm/llms/clarifai/chat/transformation.py b/litellm/llms/clarifai/chat/transformation.py index cea6e6c7440..8c2fd9e0aff 100644 --- a/litellm/llms/clarifai/chat/transformation.py +++ b/litellm/llms/clarifai/chat/transformation.py @@ -76,6 +76,13 @@ class ClarifaiConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("CLARIFAI_API_KEY") or "" return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def transform_request(self, model, messages, optional_params, litellm_params, headers): model = self.get_base_model(model) or model return super().transform_request(model, messages, optional_params, litellm_params, headers) diff --git a/litellm/llms/codestral/completion/transformation.py b/litellm/llms/codestral/completion/transformation.py index 7d675b12f4e..636760b093d 100644 --- a/litellm/llms/codestral/completion/transformation.py +++ b/litellm/llms/codestral/completion/transformation.py @@ -123,3 +123,9 @@ class CodestralTextCompletionConfig(OpenAITextCompletionConfig): finish_reason=finish_reason, logprobs=logprobs, ) + + def chunk_parser( + self, + chunk_data: str, + ) -> GenericStreamingChunk: + return self._chunk_parser(chunk_data) diff --git a/litellm/llms/cohere/embed/handler.py b/litellm/llms/cohere/embed/handler.py index 6d9543d7c30..2f60d94c08a 100644 --- a/litellm/llms/cohere/embed/handler.py +++ b/litellm/llms/cohere/embed/handler.py @@ -103,7 +103,7 @@ async def async_embedding( raise e ## PROCESS RESPONSE ## - return CohereEmbeddingConfig()._transform_response( + return CohereEmbeddingConfig().transform_response( response=response, api_key=api_key, logging_obj=logging_obj, @@ -168,7 +168,7 @@ def embedding( response: Final = client.post(embed_url, headers=headers, data=json.dumps(data)) - return CohereEmbeddingConfig()._transform_response( + return CohereEmbeddingConfig().transform_response( response=response, api_key=api_key, logging_obj=logging_obj, diff --git a/litellm/llms/cohere/embed/v1_transformation.py b/litellm/llms/cohere/embed/v1_transformation.py index b35fae5a1ac..ad5bfa507e5 100644 --- a/litellm/llms/cohere/embed/v1_transformation.py +++ b/litellm/llms/cohere/embed/v1_transformation.py @@ -68,6 +68,14 @@ class CohereEmbeddingConfig: return transformed_request + def transform_request( + self, + model: str, + input: list[str], # mutable-ok: mirrors override contract + inference_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> CohereEmbeddingRequestWithModel: + return self._transform_request(model, input, inference_params) + def _calculate_usage(self, input: list[str], encoding: _SupportsEncode, meta: dict) -> Usage: input_tokens = 0 @@ -123,6 +131,19 @@ class CohereEmbeddingConfig: input=input, ) + def transform_response( + self, + response: httpx.Response, + api_key: str | None, + logging_obj: LiteLLMLoggingObj, + data: dict[str, object] | CohereEmbeddingRequest, # mutable-ok: mirrors override contract + model_response: EmbeddingResponse, + model: str, + encoding: _SupportsEncode, + input: list[str], # mutable-ok: mirrors override contract + ) -> EmbeddingResponse: + return self._transform_response(response, api_key, logging_obj, data, model_response, model, encoding, input) + def _populate_embedding_response( self, response_json: dict, @@ -178,3 +199,13 @@ class CohereEmbeddingConfig: ) return model_response + + def populate_embedding_response( + self, + response_json: dict[str, object], # mutable-ok: mirrors override contract + model_response: EmbeddingResponse, + model: str, + encoding: _SupportsEncode, + input: list[str], # mutable-ok: mirrors override contract + ) -> EmbeddingResponse: + return self._populate_embedding_response(response_json, model_response, model, encoding, input) diff --git a/litellm/llms/custom_httpx/aiohttp_handler.py b/litellm/llms/custom_httpx/aiohttp_handler.py index 034c9514092..c3b63e207d2 100644 --- a/litellm/llms/custom_httpx/aiohttp_handler.py +++ b/litellm/llms/custom_httpx/aiohttp_handler.py @@ -18,7 +18,7 @@ from litellm.llms.custom_httpx.aiohttp_transport import LiteLLMAiohttpTransport from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, + get_httpx_client, get_ssl_configuration, ) from litellm.types.llms.openai import FileTypes @@ -60,7 +60,7 @@ class BaseLLMAIOHTTPHandler: # Create a transport using AsyncHTTPHandler's logic try: ssl_config: Final = get_ssl_configuration() - self.transport = AsyncHTTPHandler._create_aiohttp_transport( + self.transport = AsyncHTTPHandler.create_aiohttp_transport( ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, ssl_context=ssl_config if isinstance(ssl_config, ssl.SSLContext) else None, ) @@ -94,7 +94,7 @@ class BaseLLMAIOHTTPHandler: transport: Final = self.transport or self._get_or_create_transport() if transport is not None and hasattr(transport, "_get_valid_client_session"): try: - return transport._get_valid_client_session() + return transport.get_valid_client_session() except RuntimeError: pass @@ -401,7 +401,7 @@ class BaseLLMAIOHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -442,7 +442,7 @@ class BaseLLMAIOHTTPHandler: client: HTTPHandler | None = None, ) -> tuple[Any, dict]: if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client stream = True @@ -625,7 +625,7 @@ class BaseLLMAIOHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index b5e2f73c4cd..f02540eb657 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -289,7 +289,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport): cls._background_close_tasks.add(task) task.add_done_callback(cls._on_close_task_done) - def _get_valid_client_session(self) -> ClientSession: + def get_valid_client_session(self) -> ClientSession: """ Helper to get a valid ClientSession for the current event loop. @@ -340,6 +340,8 @@ class LiteLLMAiohttpTransport(AiohttpTransport): return self.client + _get_valid_client_session = get_valid_client_session + async def _make_aiohttp_request( self, client_session: ClientSession, @@ -407,7 +409,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport): sni_hostname: Final[str | None] = request.extensions.get("sni_hostname") # Use helper to ensure we have a valid session for the current event loop - client_session = self._get_valid_client_session() + client_session = self.get_valid_client_session() # Resolve proxy settings from environment variables proxy: Final = await self._get_proxy_settings(request) diff --git a/litellm/llms/custom_httpx/container_handler.py b/litellm/llms/custom_httpx/container_handler.py index 6c95816846d..5f2727e3e5a 100644 --- a/litellm/llms/custom_httpx/container_handler.py +++ b/litellm/llms/custom_httpx/container_handler.py @@ -19,8 +19,8 @@ from litellm.litellm_core_utils.url_utils import encode_url_path_segment from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.containers.main import ( ContainerFileListResponse, @@ -255,7 +255,7 @@ def _sync_http_client( ) -> HTTPHandler: """The sync HTTP client for a container request, reusing the caller's when usable.""" if client is None or not isinstance(client, HTTPHandler): - return _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + return get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) return client diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 8235a116d70..51f9f1e9fa4 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -104,7 +104,7 @@ class _TCPConnectorKwargs(TypedDict, total=False): socket_factory: Callable[[_AddrInfo], socket.socket] -def _build_aiohttp_keepalive_socket_factory() -> Callable[[_AddrInfo], socket.socket] | None: +def build_aiohttp_keepalive_socket_factory() -> Callable[[_AddrInfo], socket.socket] | None: """ Build a socket_factory that enables SO_KEEPALIVE on aiohttp TCP sockets. @@ -139,6 +139,9 @@ def _build_aiohttp_keepalive_socket_factory() -> Callable[[_AddrInfo], socket.so return factory +_build_aiohttp_keepalive_socket_factory = build_aiohttp_keepalive_socket_factory + + def get_default_headers() -> dict: """ Get default headers for HTTP requests. @@ -676,7 +679,7 @@ class AsyncHTTPHandler: timeout = _DEFAULT_TIMEOUT # Create a client with a connection pool - transport: Final = AsyncHTTPHandler._create_async_transport( + transport: Final = AsyncHTTPHandler.create_async_transport( ssl_context=ssl_config if isinstance(ssl_config, ssl.SSLContext) else None, ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, shared_session=shared_session, @@ -687,7 +690,7 @@ class AsyncHTTPHandler: return httpx.AsyncClient( transport=transport, - mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=cert), + mounts=AsyncHTTPHandler.create_httpx_proxy_mounts(transport, verify=ssl_config, cert=cert), event_hooks=event_hooks, timeout=timeout, verify=ssl_config, @@ -1175,7 +1178,7 @@ class AsyncHTTPHandler: pass @staticmethod - def _create_async_transport( + def create_async_transport( ssl_context: ssl.SSLContext | None = None, ssl_verify: bool | None = None, shared_session: Optional["ClientSession"] = None, @@ -1197,8 +1200,8 @@ class AsyncHTTPHandler: ######################################################### # AIOHTTP TRANSPORT is off by default ######################################################### - if AsyncHTTPHandler._should_use_aiohttp_transport(): - return AsyncHTTPHandler._create_aiohttp_transport( + if AsyncHTTPHandler.should_use_aiohttp_transport(): + return AsyncHTTPHandler.create_aiohttp_transport( ssl_context=ssl_context, ssl_verify=ssl_verify, shared_session=shared_session, @@ -1209,8 +1212,10 @@ class AsyncHTTPHandler: ######################################################### return AsyncHTTPHandler._create_httpx_transport() + _create_async_transport = create_async_transport + @staticmethod - def _should_use_aiohttp_transport() -> bool: + def should_use_aiohttp_transport() -> bool: """ AiohttpTransport is the default transport for litellm. @@ -1241,6 +1246,8 @@ class AsyncHTTPHandler: verbose_logger.debug("Using AiohttpTransport...") return True + _should_use_aiohttp_transport = should_use_aiohttp_transport + @staticmethod def _get_ssl_connector_kwargs( ssl_verify: bool | None = None, @@ -1270,7 +1277,7 @@ class AsyncHTTPHandler: return connector_kwargs @staticmethod - def _create_aiohttp_transport( + def create_aiohttp_transport( ssl_verify: bool | None = None, ssl_context: ssl.SSLContext | None = None, shared_session: Optional["ClientSession"] = None, @@ -1319,7 +1326,7 @@ class AsyncHTTPHandler: transport_connector_kwargs["limit_per_host"] = AIOHTTP_CONNECTOR_LIMIT_PER_HOST # Returns None when SO_KEEPALIVE is disabled or aiohttp is too old to # accept socket_factory — version detection lives inside the builder. - socket_factory: Final = _build_aiohttp_keepalive_socket_factory() + socket_factory: Final = build_aiohttp_keepalive_socket_factory() if socket_factory is not None: transport_connector_kwargs["socket_factory"] = socket_factory @@ -1347,6 +1354,8 @@ class AsyncHTTPHandler: ssl_verify=ssl_for_transport, ) + _create_aiohttp_transport = create_aiohttp_transport + @staticmethod def _create_httpx_transport() -> AsyncHTTPTransport | None: """ @@ -1361,7 +1370,7 @@ class AsyncHTTPHandler: return None @staticmethod - def _create_httpx_proxy_mounts( + def create_httpx_proxy_mounts( transport: LiteLLMAiohttpTransport | AsyncHTTPTransport | None, verify: VerifyTypes, cert: CertTypes | None, @@ -1372,6 +1381,8 @@ class AsyncHTTPHandler: lambda proxy_url: AsyncHTTPTransport(proxy=proxy_url, verify=verify, cert=cert, http2=http2_enabled()) ) + _create_httpx_proxy_mounts = create_httpx_proxy_mounts + class HTTPHandler: def __init__( @@ -1766,7 +1777,7 @@ def get_async_httpx_client( return _new_client -def _get_httpx_client(params: dict | None = None) -> HTTPHandler: +def get_httpx_client(params: dict | None = None) -> HTTPHandler: """ Retrieves the HTTP client from the cache If not present, creates a new client @@ -1810,3 +1821,6 @@ def _get_httpx_client(params: dict | None = None) -> HTTPHandler: litellm_owned_client=True, ) return _new_client + + +_get_httpx_client = get_httpx_client diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 202bfef0d8a..2b86b9bbe5c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -105,8 +105,8 @@ from litellm.llms.custom_httpx.container_handler import raise_for_error_status from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, @@ -908,7 +908,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client( + sync_httpx_client = get_httpx_client( params={"ssl_verify": litellm_params.get("ssl_verify", None)}, ) else: @@ -958,7 +958,7 @@ class BaseLLMHTTPHandler: json_mode: bool = False, ) -> tuple[object, dict]: if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client( + sync_httpx_client = get_httpx_client( { "ssl_verify": litellm_params.get("ssl_verify", None), } @@ -1262,7 +1262,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -1417,7 +1417,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -1639,7 +1639,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client() + client = get_httpx_client() json_data: Final = data if files is None and isinstance(data, dict) else None @@ -1812,7 +1812,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client() + client = get_httpx_client() # Check HTTP method from provider config http_method: Final = provider_config.get_http_method() @@ -2493,7 +2493,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -2970,7 +2970,7 @@ class BaseLLMHTTPHandler: shared_session=shared_session, ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -3061,7 +3061,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -3223,7 +3223,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -3439,7 +3439,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -3906,7 +3906,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -3994,7 +3994,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -4282,7 +4282,7 @@ class BaseLLMHTTPHandler: shared_session=shared_session, ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -4440,7 +4440,7 @@ class BaseLLMHTTPHandler: shared_session=shared_session, ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -4624,7 +4624,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -4751,7 +4751,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -4876,7 +4876,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -5085,7 +5085,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client() + sync_httpx_client = get_httpx_client() else: sync_httpx_client = client @@ -6211,6 +6211,32 @@ class BaseLLMHTTPHandler: provider_error.status_code_is_synthesized = True raise provider_error + def handle_error( + self, + e: Exception, + provider_config: Union[ + BaseConfig, + BaseRerankConfig, + BaseResponsesAPIConfig, + BaseImageEditConfig, + BaseImageGenerationConfig, + BaseVectorStoreConfig, + BaseVectorStoreFilesConfig, + BaseGoogleGenAIGenerateContentConfig, + BaseAnthropicMessagesConfig, + BaseBatchesConfig, + BaseVideoConfig, + BaseSearchConfig, + BaseTextToSpeechConfig, + BaseSkillsAPIConfig, + "BasePassthroughConfig", + "BaseContainerConfig", + BaseEvalsAPIConfig, + BaseRealtimeHTTPConfig, + ], + ) -> None: + return self._handle_error(e, provider_config) + @staticmethod def _append_query_params(url: str, query_params: RealtimeQueryParams | None) -> str: """Append query_params to url, skipping keys already present in the URL.""" @@ -6841,7 +6867,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -7060,7 +7086,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -7294,7 +7320,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -7516,7 +7542,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -7691,7 +7717,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client directly (like video_generation does) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -7865,7 +7891,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -8019,7 +8045,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -8152,7 +8178,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -8360,7 +8386,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -8740,7 +8766,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client directly (like video_generation does) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -8931,7 +8957,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -9102,7 +9128,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -9269,7 +9295,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -9438,7 +9464,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -9612,7 +9638,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -9781,7 +9807,7 @@ class BaseLLMHTTPHandler: # For sync calls, use sync HTTP client if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -10117,7 +10143,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -10293,7 +10319,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -10419,7 +10445,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -10555,7 +10581,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -10697,7 +10723,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -10825,7 +10851,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -10958,7 +10984,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -11100,7 +11126,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -11235,7 +11261,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -11368,7 +11394,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -11511,7 +11537,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -11645,7 +11671,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -11737,7 +11763,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -11978,7 +12004,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12223,7 +12249,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12345,7 +12371,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12451,7 +12477,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12554,7 +12580,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12661,7 +12687,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12767,7 +12793,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12873,7 +12899,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -12978,7 +13004,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -13082,7 +13108,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -13183,7 +13209,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -13290,7 +13316,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -13396,7 +13422,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -13502,7 +13528,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -13605,7 +13631,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client @@ -13706,7 +13732,7 @@ class BaseLLMHTTPHandler: ) if client is None or not isinstance(client, HTTPHandler): - sync_httpx_client = _get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) + sync_httpx_client = get_httpx_client(params={"ssl_verify": litellm_params.get("ssl_verify", None)}) else: sync_httpx_client = client diff --git a/litellm/llms/dashscope/chat/transformation.py b/litellm/llms/dashscope/chat/transformation.py index 5ef482a6c2e..867a183dd07 100644 --- a/litellm/llms/dashscope/chat/transformation.py +++ b/litellm/llms/dashscope/chat/transformation.py @@ -60,6 +60,13 @@ class DashScopeChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("DASHSCOPE_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def _resolve_chat_api_base(self, api_base: str | None) -> str: return api_base or "https://dashscope.aliyuncs.com/compatible-mode/v1" diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 5fd553126cc..89b614d6258 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -729,7 +729,7 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator): # 6. Set tool_calls to None from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.llms.base_llm.base_utils import ( - _convert_tool_response_to_message, + convert_tool_response_to_message, ) # Check if this chunk has a function name @@ -744,7 +744,7 @@ class DatabricksChatResponseIterator(BaseModelResponseIterator): or function_name == RESPONSE_FORMAT_TOOL_NAME ): # Convert tool calls to message format - message = _convert_tool_response_to_message(tool_calls) + message = convert_tool_response_to_message(tool_calls) if message is not None: if message.content == "{}": # empty json message.content = "" diff --git a/litellm/llms/datarobot/chat/transformation.py b/litellm/llms/datarobot/chat/transformation.py index ee787263c3f..735ed1ec0b5 100644 --- a/litellm/llms/datarobot/chat/transformation.py +++ b/litellm/llms/datarobot/chat/transformation.py @@ -69,6 +69,13 @@ class DataRobotConfig(OpenAILikeChatConfig): return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_complete_url( self, api_base: str | None, diff --git a/litellm/llms/deepinfra/chat/transformation.py b/litellm/llms/deepinfra/chat/transformation.py index e746ce0a5c9..004128c0184 100644 --- a/litellm/llms/deepinfra/chat/transformation.py +++ b/litellm/llms/deepinfra/chat/transformation.py @@ -197,3 +197,10 @@ class DeepInfraConfig(OpenAIGPTConfig): api_base = api_base or get_secret_str("DEEPINFRA_API_BASE") or "https://api.deepinfra.com/v1/openai" dynamic_api_key: Final = api_key or get_secret_str("DEEPINFRA_API_KEY") return api_base, dynamic_api_key + + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) diff --git a/litellm/llms/docker_model_runner/chat/transformation.py b/litellm/llms/docker_model_runner/chat/transformation.py index 303710a7108..adb40311efd 100644 --- a/litellm/llms/docker_model_runner/chat/transformation.py +++ b/litellm/llms/docker_model_runner/chat/transformation.py @@ -65,6 +65,13 @@ class DockerModelRunnerChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("DOCKER_MODEL_RUNNER_API_KEY") or "dummy-key" return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_complete_url( self, api_base: str | None, diff --git a/litellm/llms/fal_ai/videos/transformation.py b/litellm/llms/fal_ai/videos/transformation.py index cbcd92acbaf..9000588307e 100644 --- a/litellm/llms/fal_ai/videos/transformation.py +++ b/litellm/llms/fal_ai/videos/transformation.py @@ -16,8 +16,8 @@ from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared HTTP factory is private get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared HTTP factory lacks typed params + get_httpx_client, # pyright: ignore[reportPrivateUsage, reportUnknownVariableType] # shared HTTP factory is private ) from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -279,7 +279,7 @@ def _status_video_object( class FalAIVideoConfig(BaseVideoConfig): def __init__( self, - sync_client_factory: Callable[[], HTTPHandler] = _get_httpx_client, + sync_client_factory: Callable[[], HTTPHandler] = get_httpx_client, async_client_factory: Callable[[], AsyncHTTPHandler] = _get_fal_ai_async_httpx_client, ) -> None: super().__init__() diff --git a/litellm/llms/featherless_ai/chat/transformation.py b/litellm/llms/featherless_ai/chat/transformation.py index d062392488b..a68261c58e9 100644 --- a/litellm/llms/featherless_ai/chat/transformation.py +++ b/litellm/llms/featherless_ai/chat/transformation.py @@ -110,6 +110,13 @@ class FeatherlessAIConfig(OpenAIGPTConfig): dynamic_api_key = api_key or get_secret_str("FEATHERLESS_AI_API_KEY") or get_secret_str("FEATHERLESS_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def validate_environment( self, headers: dict, diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 0a8a010f2b9..a7a417db2c9 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -787,6 +787,13 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): ) return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_models(self, api_key: str | None = None, api_base: str | None = None): api_base, api_key = self._get_openai_compatible_provider_info(api_base=api_base, api_key=api_key) if api_base is None or api_key is None: diff --git a/litellm/llms/fireworks_ai/completion/transformation.py b/litellm/llms/fireworks_ai/completion/transformation.py index e207ae0cecf..1f5d4f306a9 100644 --- a/litellm/llms/fireworks_ai/completion/transformation.py +++ b/litellm/llms/fireworks_ai/completion/transformation.py @@ -6,7 +6,7 @@ from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUser from litellm.utils import supports_reasoning from ...base_llm.completion.transformation import BaseTextCompletionConfig -from ...openai.completion.utils import _transform_prompt +from ...openai.completion.utils import transform_prompt from ..chat.transformation import ( EFFORT_KWARG_KEYS, NIM_VLLM_STRIP_PARAMS, @@ -159,7 +159,7 @@ class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig headers: dict, ) -> dict: translated_params: Final = self.map_extra_body_params(optional_params=optional_params, model=model) - prompt: Final = _transform_prompt(messages=messages) + prompt: Final = transform_prompt(messages=messages) data: Final = { "model": resolve_fireworks_resource_name(model), diff --git a/litellm/llms/gemini/chat/transformation.py b/litellm/llms/gemini/chat/transformation.py index f833837f97c..143652ee77c 100644 --- a/litellm/llms/gemini/chat/transformation.py +++ b/litellm/llms/gemini/chat/transformation.py @@ -13,7 +13,7 @@ from litellm.types.llms.openai import AllMessageValues, ChatCompletionFileObject from litellm.types.llms.vertex_ai import ContentType, PartType from litellm.utils import supports_reasoning -from ...vertex_ai.gemini.transformation import GEMINI_FILES_API_URI_PREFIX, _gemini_convert_messages_with_history +from ...vertex_ai.gemini.transformation import GEMINI_FILES_API_URI_PREFIX, gemini_convert_messages_with_history from ...vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig @@ -166,9 +166,17 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig): except Exception: # If conversion fails, leave as is and let the API handle it pass - return _gemini_convert_messages_with_history( + return gemini_convert_messages_with_history( messages=messages, model=model, litellm_params=litellm_params, custom_llm_provider="gemini", ) + + def transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: mirrors override contract + model: str | None = None, + litellm_params: dict[str, object] | None = None, # mutable-ok: mirrors override contract + ) -> list[ContentType]: # mutable-ok: mirrors override contract + return self._transform_messages(messages, model, litellm_params) diff --git a/litellm/llms/gemini/common_utils.py b/litellm/llms/gemini/common_utils.py index c44db8cac26..92b3b85dc51 100644 --- a/litellm/llms/gemini/common_utils.py +++ b/litellm/llms/gemini/common_utils.py @@ -192,7 +192,7 @@ def _dedupe_gemini_search_tools(tools: list[dict[str, object]]) -> list[dict[str VertexGeminiConfig, ) - search_tool_keys: Final = VertexGeminiConfig._search_tool_keys() + search_tool_keys: Final = VertexGeminiConfig.search_tool_keys() seen_search_keys: Final[set[str]] = set() deduped_tools: Final[list[dict[str, object]]] = [] @@ -220,7 +220,7 @@ def _has_gemini_search_tool(tools: list[object]) -> bool: VertexGeminiConfig, ) - search_tool_keys: Final = VertexGeminiConfig._search_tool_keys() + search_tool_keys: Final = VertexGeminiConfig.search_tool_keys() return any(isinstance(tool, dict) and any(key in tool for key in search_tool_keys) for tool in tools) @@ -238,18 +238,18 @@ def map_gemini_image_tools_params( tools_value: Final = non_default_params.get("tools") if isinstance(tools_value, list) and tools_value: - mapped_tools: Final = gemini_config._map_function(value=tools_value, optional_params=result) - result = gemini_config._add_tools_to_optional_params(result, mapped_tools) + mapped_tools: Final = gemini_config.map_function(value=tools_value, optional_params=result) + result = gemini_config.add_tools_to_optional_params(result, mapped_tools) web_search_options: Final = non_default_params.get("web_search_options") existing_tools: Final = result.get("tools") if isinstance(web_search_options, dict) and not ( isinstance(existing_tools, list) and _has_gemini_search_tool(existing_tools) ): - search_tool: Final = gemini_config._map_web_search_options(web_search_options) - result = gemini_config._add_tools_to_optional_params(result, [search_tool]) + search_tool: Final = gemini_config.map_web_search_options(web_search_options) + result = gemini_config.add_tools_to_optional_params(result, [search_tool]) - gemini_config._drop_search_tools_mixed_with_functions(result) + gemini_config.drop_search_tools_mixed_with_functions(result) resolved_tools: Final = result.get("tools") if isinstance(resolved_tools, list): @@ -277,7 +277,7 @@ def get_gemini_image_web_search_requests( elif isinstance(candidate_grounding, dict): grounding_metadata.append(candidate_grounding) - return VertexGeminiConfig._calculate_web_search_requests(grounding_metadata) + return VertexGeminiConfig.calculate_web_search_requests(grounding_metadata) def get_gemini_image_generation_config( diff --git a/litellm/llms/gemini/google_genai/transformation.py b/litellm/llms/gemini/google_genai/transformation.py index 1189af2d6a3..2eb84423633 100644 --- a/litellm/llms/gemini/google_genai/transformation.py +++ b/litellm/llms/gemini/google_genai/transformation.py @@ -13,7 +13,7 @@ from litellm.llms.base_llm.google_genai.transformation import ( BaseGoogleGenAIGenerateContentConfig, ) from litellm.llms.vertex_ai.common_utils import ( - _build_vertex_schema, + build_vertex_schema, supports_response_json_schema, ) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM @@ -113,22 +113,22 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): Mapped parameters for the provider """ from litellm.llms.vertex_ai.gemini.transformation import ( - _camel_to_snake, - _snake_to_camel, + camel_to_snake, + snake_to_camel, ) _generate_content_config_dict: Final[dict[str, object]] = {} supported_google_genai_params: Final = self.get_supported_generate_content_optional_params(model) # Create a set with both camelCase and snake_case versions for faster lookup supported_params_set: Final = set(supported_google_genai_params) - supported_params_set.update(_snake_to_camel(p) for p in supported_google_genai_params) - supported_params_set.update(_camel_to_snake(p) for p in supported_google_genai_params if "_" not in p) + supported_params_set.update(snake_to_camel(p) for p in supported_google_genai_params) + supported_params_set.update(camel_to_snake(p) for p in supported_google_genai_params if "_" not in p) for param, value in generate_content_config_dict.items(): # Google GenAI API expects camelCase, so we'll always output in camelCase # Check if param (or its variants) is supported - param_snake = _camel_to_snake(param) - param_camel = _snake_to_camel(param) + param_snake = camel_to_snake(param) + param_camel = snake_to_camel(param) # Check if param is supported in any format is_supported = ( @@ -327,7 +327,7 @@ class GoogleGenAIConfig(BaseGoogleGenAIGenerateContentConfig, VertexLLM): else: if json_schema_key is not None: generate_content_config_dict.pop(json_schema_key) - generate_content_config_dict[schema_key] = _build_vertex_schema( + generate_content_config_dict[schema_key] = build_vertex_schema( parameters=deepcopy(value), add_property_ordering=True ) diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 3b926ccd1c2..7c04ebf311f 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -95,7 +95,7 @@ def _gemini_live_speech_config(voice: object) -> Mapping[str, object] | None: voice, ) return None - return VertexGeminiConfig()._map_audio_params({"voice": voice}) + return VertexGeminiConfig().map_audio_params({"voice": voice}) class _GeminiLiveSetupEnvelope(TypedDict, total=False): @@ -331,7 +331,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): vertex_gemini_config = VertexGeminiConfig() # Tools should be at the top level of setup, not inside generationConfig - optional_params["tools"] = vertex_gemini_config._map_function( + optional_params["tools"] = vertex_gemini_config.map_function( value=value, optional_params=optional_params ) elif key == "input_audio_transcription" and value is not None: @@ -1056,7 +1056,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): _modalities: Final = [modality.lower() for modality in cast(list[str], gemini_modalities)] resolved_usage_metadata: Final = self._consume_usage_metadata_for_response_done(cast(dict, message)) if resolved_usage_metadata is not None: - _chat_completion_usage = VertexGeminiConfig._calculate_usage( + _chat_completion_usage = VertexGeminiConfig.calculate_usage( completion_response=cast( BidiGenerateContentServerMessage, {**cast(dict, message), "usageMetadata": resolved_usage_metadata}, @@ -1484,7 +1484,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): resolved_tool_call_usage_metadata = self._consume_usage_metadata_for_response_done(json_message) if resolved_tool_call_usage_metadata is not None: - _tool_call_chat_completion_usage = VertexGeminiConfig._calculate_usage( + _tool_call_chat_completion_usage = VertexGeminiConfig.calculate_usage( completion_response=cast( BidiGenerateContentServerMessage, { diff --git a/litellm/llms/gigachat/authenticator.py b/litellm/llms/gigachat/authenticator.py index a85dcd9c70d..d43e203d7a1 100644 --- a/litellm/llms/gigachat/authenticator.py +++ b/litellm/llms/gigachat/authenticator.py @@ -18,8 +18,8 @@ from litellm.caching.caching import InMemoryCache from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, - _get_httpx_client, # pyright: ignore[reportPrivateUsage] # house cached-client factory has no public alias get_async_httpx_client, + get_httpx_client, ) from litellm.secret_managers.main import get_secret_str from litellm.types.utils import LlmProviders @@ -58,7 +58,7 @@ def _get_scope() -> str: def _get_http_client() -> HTTPHandler: """Get cached httpx client with SSL verification disabled.""" - return _get_httpx_client(params={"ssl_verify": False}) + return get_httpx_client(params={"ssl_verify": False}) def get_access_token( diff --git a/litellm/llms/gigachat/file_handler.py b/litellm/llms/gigachat/file_handler.py index 359553e144f..c051d693720 100644 --- a/litellm/llms/gigachat/file_handler.py +++ b/litellm/llms/gigachat/file_handler.py @@ -14,8 +14,8 @@ from typing import Final from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.llms.gigachat.utils import get_api_base from litellm.types.utils import LlmProviders @@ -60,7 +60,7 @@ def _content_type_or_default(headers: Mapping[str, str]) -> str: def _download_image_sync(url: str) -> tuple[bytes, str, str]: """Download image from URL synchronously.""" - client: Final = _get_httpx_client(params={"ssl_verify": False}) + client: Final = get_httpx_client(params={"ssl_verify": False}) response: Final = client.get(url) response.raise_for_status() @@ -128,7 +128,7 @@ def upload_file_sync( base_url: Final = get_api_base(api_base) upload_url: Final = f"{base_url}/files" - client: Final = _get_httpx_client(params={"ssl_verify": False}) + client: Final = get_httpx_client(params={"ssl_verify": False}) response: Final = client.post( upload_url, headers={"Authorization": f"Bearer {access_token}"}, diff --git a/litellm/llms/github_copilot/authenticator.py b/litellm/llms/github_copilot/authenticator.py index 83867708c46..a58f6b08a89 100644 --- a/litellm/llms/github_copilot/authenticator.py +++ b/litellm/llms/github_copilot/authenticator.py @@ -10,7 +10,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm._logging import verbose_logger from litellm.litellm_core_utils.asyncify import can_block_current_thread -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client from .common_utils import ( APIKeyExpiredError, @@ -195,7 +195,7 @@ class Authenticator: max_retries: Final = 3 for attempt in range(max_retries): try: - sync_client = _get_httpx_client() + sync_client = get_httpx_client() response = sync_client.get(api_key_url, headers=headers) response.raise_for_status() @@ -257,7 +257,7 @@ class Authenticator: GetDeviceCodeError: If unable to get a device code. """ try: - sync_client: Final = _get_httpx_client() + sync_client: Final = get_httpx_client() device_code_url: Final = os.getenv("GITHUB_COPILOT_DEVICE_CODE_URL", DEFAULT_GITHUB_DEVICE_CODE_URL) client_id: Final = os.getenv("GITHUB_COPILOT_CLIENT_ID", DEFAULT_GITHUB_CLIENT_ID) resp: Final = sync_client.post( @@ -309,7 +309,7 @@ class Authenticator: Raises: GetAccessTokenError: If unable to get an access token. """ - sync_client: Final = _get_httpx_client() + sync_client: Final = get_httpx_client() max_attempts: Final = 12 # 1 minute (12 * 5 seconds) access_token_url: Final = os.getenv("GITHUB_COPILOT_ACCESS_TOKEN_URL", DEFAULT_GITHUB_ACCESS_TOKEN_URL) diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 568cd4365cc..b17bd74365c 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -58,6 +58,15 @@ class GithubCopilotConfig(OpenAIConfig): ) return dynamic_api_base, dynamic_api_key, custom_llm_provider + def get_openai_compatible_provider_info( + self, + model: str, + api_base: str | None, + api_key: str | None, + custom_llm_provider: str, + ) -> tuple[str | None, str | None, str]: + return self._get_openai_compatible_provider_info(model, api_base, api_key, custom_llm_provider) + def _transform_messages( self, messages, diff --git a/litellm/llms/gradient_ai/chat/transformation.py b/litellm/llms/gradient_ai/chat/transformation.py index 3fc7fbcef08..cfd08e6cc95 100644 --- a/litellm/llms/gradient_ai/chat/transformation.py +++ b/litellm/llms/gradient_ai/chat/transformation.py @@ -126,6 +126,13 @@ class GradientAIConfig(OpenAILikeChatConfig): dynamic_api_key: Final = api_key or get_secret_str("GRADIENT_AI_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def map_openai_params( self, non_default_params: dict, diff --git a/litellm/llms/groq/chat/handler.py b/litellm/llms/groq/chat/handler.py index bd38e4c14e9..1346378a8e8 100644 --- a/litellm/llms/groq/chat/handler.py +++ b/litellm/llms/groq/chat/handler.py @@ -44,7 +44,7 @@ class GroqChatCompletion(OpenAILikeChatHandler): streaming_decoder: CustomStreamingDecoder | None = None, fake_stream: bool = False, ): - messages = GroqChatConfig()._transform_messages(messages=cast(list[AllMessageValues], messages), model=model) + messages = GroqChatConfig().transform_messages(messages=cast(list[AllMessageValues], messages), model=model) if optional_params.get("stream") is True: fake_stream = GroqChatConfig()._should_fake_stream(optional_params) diff --git a/litellm/llms/groq/chat/transformation.py b/litellm/llms/groq/chat/transformation.py index e3a60acebe2..ebe3319b149 100644 --- a/litellm/llms/groq/chat/transformation.py +++ b/litellm/llms/groq/chat/transformation.py @@ -159,6 +159,17 @@ class GroqChatConfig(OpenAILikeChatConfig): else: return super()._transform_messages(messages=messages, model=model, is_async=False) + def transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: mirrors override contract + model: str, + is_async: bool = False, + ) -> ( + list[AllMessageValues] # mutable-ok: mirrors override contract + | Coroutine[object, object, list[AllMessageValues]] + ): + return self._transform_messages(messages, model, is_async) + def _get_openai_compatible_provider_info( self, api_base: str | None, api_key: str | None ) -> tuple[str | None, str | None]: @@ -167,6 +178,13 @@ class GroqChatConfig(OpenAILikeChatConfig): dynamic_api_key: Final = api_key or get_secret_str("GROQ_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def _should_fake_stream(self, optional_params: dict) -> bool: """ Groq doesn't support 'response_format' while streaming diff --git a/litellm/llms/heroku/chat/transformation.py b/litellm/llms/heroku/chat/transformation.py index 23b458304ab..f41b25a4d69 100644 --- a/litellm/llms/heroku/chat/transformation.py +++ b/litellm/llms/heroku/chat/transformation.py @@ -55,6 +55,13 @@ class HerokuChatConfig(OpenAIGPTConfig): return api_base, api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_complete_url( self, api_base: str | None, diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index 58ab297a198..f3f6bc9cbf1 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -112,6 +112,13 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("HOSTED_VLLM_API_KEY") or "fake-api-key" return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def _is_video_file(self, content_item: ChatCompletionFileObject) -> bool: file: Final = content_item.get("file", {}) format: Final = file.get("format") diff --git a/litellm/llms/huggingface/chat/transformation.py b/litellm/llms/huggingface/chat/transformation.py index 0b633d67eda..b45c706f8bf 100644 --- a/litellm/llms/huggingface/chat/transformation.py +++ b/litellm/llms/huggingface/chat/transformation.py @@ -16,7 +16,7 @@ else: from litellm.llms.base_llm.chat.transformation import BaseLLMException from ...openai.chat.gpt_transformation import OpenAIGPTConfig -from ..common_utils import HuggingFaceError, _fetch_inference_provider_mapping +from ..common_utils import HuggingFaceError, fetch_inference_provider_mapping logger: Final = logging.getLogger(__name__) @@ -138,7 +138,7 @@ class HuggingFaceChatConfig(OpenAIGPTConfig): if "/" in remaining: provider: Final = first_part model_id: Final = remaining - provider_mapping = _fetch_inference_provider_mapping(model_id) + provider_mapping = fetch_inference_provider_mapping(model_id) if provider not in provider_mapping: raise HuggingFaceError( message=f"Model {model_id} is not supported for provider {provider}", diff --git a/litellm/llms/huggingface/common_utils.py b/litellm/llms/huggingface/common_utils.py index 30da3ce091d..08d89b50156 100644 --- a/litellm/llms/huggingface/common_utils.py +++ b/litellm/llms/huggingface/common_utils.py @@ -58,7 +58,7 @@ def output_parser(generated_text: str): @lru_cache(maxsize=128) -def _fetch_inference_provider_mapping(model: str) -> dict: +def fetch_inference_provider_mapping(model: str) -> dict: """ Fetch provider mappings for a model from the Hugging Face Hub. @@ -100,3 +100,6 @@ def _fetch_inference_provider_mapping(model: str) -> dict: status_code=status_code, headers=headers, ) + + +_fetch_inference_provider_mapping = fetch_inference_provider_mapping diff --git a/litellm/llms/hyperbolic/chat/transformation.py b/litellm/llms/hyperbolic/chat/transformation.py index 9ec95e7a9d5..7ad750df040 100644 --- a/litellm/llms/hyperbolic/chat/transformation.py +++ b/litellm/llms/hyperbolic/chat/transformation.py @@ -30,6 +30,13 @@ class HyperbolicChatConfig(OpenAILikeChatConfig): dynamic_api_key: Final = api_key or get_secret_str("HYPERBOLIC_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_supported_openai_params(self, model: str) -> list: """ Hyperbolic supports standard OpenAI parameters diff --git a/litellm/llms/inception/chat/transformation.py b/litellm/llms/inception/chat/transformation.py index 0af9e06c10d..31675f6d650 100644 --- a/litellm/llms/inception/chat/transformation.py +++ b/litellm/llms/inception/chat/transformation.py @@ -50,3 +50,10 @@ class InceptionChatConfig(OpenAILikeChatConfig): if passed_api_base is None or api_key: dynamic_api_key = api_key or litellm.inception_key or get_secret_str("INCEPTION_API_KEY") return api_base, dynamic_api_key + + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) diff --git a/litellm/llms/jina_ai/embedding/transformation.py b/litellm/llms/jina_ai/embedding/transformation.py index e184ad2628b..683b73ce09b 100644 --- a/litellm/llms/jina_ai/embedding/transformation.py +++ b/litellm/llms/jina_ai/embedding/transformation.py @@ -90,6 +90,13 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): ) return LlmProviders.JINA_AI.value, api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_complete_url( self, api_base: str | None, diff --git a/litellm/llms/lambda_ai/chat/transformation.py b/litellm/llms/lambda_ai/chat/transformation.py index fedce35cd28..dc517a09c1b 100644 --- a/litellm/llms/lambda_ai/chat/transformation.py +++ b/litellm/llms/lambda_ai/chat/transformation.py @@ -27,3 +27,10 @@ class LambdaAIChatConfig(OpenAILikeChatConfig): ) dynamic_api_key: Final = api_key or get_secret_str("LAMBDA_API_KEY") return api_base, dynamic_api_key + + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) diff --git a/litellm/llms/langflow/chat/transformation.py b/litellm/llms/langflow/chat/transformation.py index c5bb293951a..1781082d36f 100644 --- a/litellm/llms/langflow/chat/transformation.py +++ b/litellm/llms/langflow/chat/transformation.py @@ -53,6 +53,13 @@ class LangFlowConfig(BaseConfig): api_key = api_key or get_secret_str("LANGFLOW_API_KEY") return api_base, api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_supported_openai_params(self, model: str) -> list[str]: return ["stream"] diff --git a/litellm/llms/langgraph/chat/transformation.py b/litellm/llms/langgraph/chat/transformation.py index 36c761e308e..f1f4054099a 100644 --- a/litellm/llms/langgraph/chat/transformation.py +++ b/litellm/llms/langgraph/chat/transformation.py @@ -71,6 +71,13 @@ class LangGraphConfig(BaseConfig): return api_base, api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_supported_openai_params(self, model: str) -> list[str]: """ LangGraph supports minimal OpenAI params since it's an agent runtime. @@ -295,12 +302,12 @@ class LangGraphConfig(BaseConfig): """ from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, - _get_httpx_client, + get_httpx_client, ) from litellm.utils import CustomStreamWrapper if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={}) + client = get_httpx_client(params={}) verbose_logger.debug("Making sync streaming request to: %s", api_base) diff --git a/litellm/llms/lemonade/chat/transformation.py b/litellm/llms/lemonade/chat/transformation.py index 341e8dd2e12..c4079dfc73b 100644 --- a/litellm/llms/lemonade/chat/transformation.py +++ b/litellm/llms/lemonade/chat/transformation.py @@ -216,6 +216,13 @@ class LemonadeChatConfig(OpenAILikeChatConfig): key = api_key or litellm.lemonade_key or get_secret_str("LEMONADE_API_KEY") or self._DEFAULT_API_KEY return api_base, key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def _get_auth_headers(self, api_key: str | None) -> dict: if api_key is None or api_key == self._DEFAULT_API_KEY: return {} diff --git a/litellm/llms/litellm_proxy/chat/transformation.py b/litellm/llms/litellm_proxy/chat/transformation.py index cf4c41cd3f3..4183f900e26 100644 --- a/litellm/llms/litellm_proxy/chat/transformation.py +++ b/litellm/llms/litellm_proxy/chat/transformation.py @@ -42,6 +42,13 @@ class LiteLLMProxyChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("LITELLM_PROXY_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: api_base, api_key = self._get_openai_compatible_provider_info(api_base, api_key) if api_base is None: diff --git a/litellm/llms/llamafile/chat/transformation.py b/litellm/llms/llamafile/chat/transformation.py index 1f51bfb0af2..013493322fa 100644 --- a/litellm/llms/llamafile/chat/transformation.py +++ b/litellm/llms/llamafile/chat/transformation.py @@ -41,3 +41,10 @@ class LlamafileChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = LlamafileChatConfig._resolve_api_key(api_key) return api_base, dynamic_api_key + + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) diff --git a/litellm/llms/lm_studio/chat/transformation.py b/litellm/llms/lm_studio/chat/transformation.py index 54a73bdc053..bebdf436a8e 100644 --- a/litellm/llms/lm_studio/chat/transformation.py +++ b/litellm/llms/lm_studio/chat/transformation.py @@ -19,6 +19,13 @@ class LMStudioChatConfig(OpenAIGPTConfig): ) # LM Studio does not require an api key, but OpenAI client requires non-None value return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def map_openai_params( self, non_default_params: dict, diff --git a/litellm/llms/mistral/chat/transformation.py b/litellm/llms/mistral/chat/transformation.py index ede9467f7f8..0d45482c101 100644 --- a/litellm/llms/mistral/chat/transformation.py +++ b/litellm/llms/mistral/chat/transformation.py @@ -286,6 +286,24 @@ class MistralConfig(OpenAIGPTConfig): else: return super()._transform_messages(new_messages, model, False) + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + + def transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: mirrors override contract + model: str, + is_async: bool = False, + ) -> ( + list[AllMessageValues] # mutable-ok: mirrors override contract + | Coroutine[object, object, list[AllMessageValues]] + ): + return self._transform_messages(messages, model, is_async) + async def _transform_messages_async(self, messages: list[AllMessageValues], model: str) -> list[AllMessageValues]: """ Handle modification of messages for Mistral API in an async context. diff --git a/litellm/llms/modelscope/chat/transformation.py b/litellm/llms/modelscope/chat/transformation.py index 27c3bff5346..0a9a652fe6a 100644 --- a/litellm/llms/modelscope/chat/transformation.py +++ b/litellm/llms/modelscope/chat/transformation.py @@ -66,6 +66,13 @@ class ModelScopeChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("MODELSCOPE_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + @override def get_complete_url( self, diff --git a/litellm/llms/moonshot/chat/transformation.py b/litellm/llms/moonshot/chat/transformation.py index 8e428d93b2d..8f06c355261 100644 --- a/litellm/llms/moonshot/chat/transformation.py +++ b/litellm/llms/moonshot/chat/transformation.py @@ -74,6 +74,13 @@ class MoonshotChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("MOONSHOT_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_complete_url( self, api_base: str | None, diff --git a/litellm/llms/morph/chat/transformation.py b/litellm/llms/morph/chat/transformation.py index 434397c8dfd..d5cba501477 100644 --- a/litellm/llms/morph/chat/transformation.py +++ b/litellm/llms/morph/chat/transformation.py @@ -30,6 +30,13 @@ class MorphChatConfig(OpenAILikeChatConfig): dynamic_api_key: Final = api_key or get_secret_str("MORPH_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_supported_openai_params(self, model: str) -> list: return [ "messages", diff --git a/litellm/llms/nlp_cloud/chat/handler.py b/litellm/llms/nlp_cloud/chat/handler.py index d0e1a8148a2..cd6be3c2224 100644 --- a/litellm/llms/nlp_cloud/chat/handler.py +++ b/litellm/llms/nlp_cloud/chat/handler.py @@ -6,7 +6,7 @@ import litellm from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, + get_httpx_client, ) from litellm.utils import ModelResponse @@ -73,7 +73,7 @@ def completion( ) ## COMPLETION CALL if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client() + client = get_httpx_client() response: Final = client.post( completion_url, diff --git a/litellm/llms/nscale/chat/transformation.py b/litellm/llms/nscale/chat/transformation.py index fc8784328a4..4f48a1fc8fd 100644 --- a/litellm/llms/nscale/chat/transformation.py +++ b/litellm/llms/nscale/chat/transformation.py @@ -33,6 +33,13 @@ class NscaleConfig(OpenAIGPTConfig): resolved_api_key: Final = NscaleConfig.get_api_key(api_key) return resolved_api_base, resolved_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_supported_openai_params(self, model: str) -> list: return [ "max_tokens", diff --git a/litellm/llms/oci/chat/cohere.py b/litellm/llms/oci/chat/cohere.py index 1b494ebad47..9163ddfa5de 100644 --- a/litellm/llms/oci/chat/cohere.py +++ b/litellm/llms/oci/chat/cohere.py @@ -15,8 +15,8 @@ import httpx from pydantic import JsonValue, TypeAdapter, ValidationError from litellm.llms.oci.chat.generic import ( - _normalize_oci_finish_reason, - _synthesize_oci_tool_call_id, + normalize_oci_finish_reason, + synthesize_oci_tool_call_id, ) from litellm.llms.oci.common_utils import ( OCI_JSON_TO_PYTHON_TYPES, @@ -76,11 +76,13 @@ def _content_text(content: str | Iterable[Mapping[str, object]] | None) -> str: return str(content) -def _extract_text_content(content: str | Iterable[Mapping[str, object]] | None) -> str: +def extract_text_content(content: str | Iterable[Mapping[str, object]] | None) -> str: """Return the plain-text representation of a message content value.""" return _content_text(content) +_extract_text_content = extract_text_content + _TOOL_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object]) @@ -138,7 +140,7 @@ def adapt_messages_to_cohere_standard( chat_history: Final[list[CohereMessage]] = [] for msg in history_source: role = msg.get("role") - content = _extract_text_content(msg.get("content")) + content = extract_text_content(msg.get("content")) tool_calls = ( [_to_cohere_tool_call(tool_call) for tool_call in msg["tool_calls"]] @@ -235,13 +237,13 @@ def handle_cohere_response( model_response.created = int(datetime.datetime.now().timestamp()) response_text: Final = cohere_response.chatResponse.text - finish_reason: Final = _normalize_oci_finish_reason(cohere_response.chatResponse.finishReason) + finish_reason: Final = normalize_oci_finish_reason(cohere_response.chatResponse.finishReason) tool_calls: list[dict[str, object]] | None = None if cohere_response.chatResponse.toolCalls: tool_calls = [ { - "id": _synthesize_oci_tool_call_id(i, tc.name, json.dumps(tc.parameters, sort_keys=True)), + "id": synthesize_oci_tool_call_id(i, tc.name, json.dumps(tc.parameters, sort_keys=True)), "type": "function", "function": { "name": tc.name, @@ -323,7 +325,7 @@ def handle_cohere_stream_chunk( # deterministically from the call's content/position. A random # uuid4 per chunk would cause downstream stream-mergers to # treat each chunk as a distinct tool call. - "id": _synthesize_oci_tool_call_id(i, tc.name, json.dumps(tc.parameters, sort_keys=True)), + "id": synthesize_oci_tool_call_id(i, tc.name, json.dumps(tc.parameters, sort_keys=True)), "type": "function", "function": { "name": tc.name, @@ -333,7 +335,7 @@ def handle_cohere_stream_chunk( for i, tc in enumerate(cohere_tool_calls) ] - finish_reason: Final = _normalize_oci_finish_reason(typed_chunk.finishReason) + finish_reason: Final = normalize_oci_finish_reason(typed_chunk.finishReason) return ModelResponseStream( choices=[ diff --git a/litellm/llms/oci/chat/generic.py b/litellm/llms/oci/chat/generic.py index 8db5803deb5..13cea58495d 100644 --- a/litellm/llms/oci/chat/generic.py +++ b/litellm/llms/oci/chat/generic.py @@ -240,7 +240,7 @@ def adapt_tool_definition_to_oci_standard(tools: list[dict], vendor: OCIVendors) return new_tools -def _normalize_oci_finish_reason(raw: str | None) -> str | None: +def normalize_oci_finish_reason(raw: str | None) -> str | None: """Map an OCI-specific finish reason to its OpenAI-standard equivalent. OCI emits ``COMPLETE`` / ``MAX_TOKENS`` / ``TOOL_CALL(S)`` plus a long tail @@ -261,7 +261,10 @@ def _normalize_oci_finish_reason(raw: str | None) -> str | None: return "stop" -def _synthesize_oci_tool_call_id(position: int, name: str, arguments: str) -> str: +_normalize_oci_finish_reason = normalize_oci_finish_reason + + +def synthesize_oci_tool_call_id(position: int, name: str, arguments: str) -> str: """Deterministic synthetic tool-call id derived from chunk content. Used as a fallback when OCI omits ``id`` (always the case for the OCI @@ -279,13 +282,16 @@ def _synthesize_oci_tool_call_id(position: int, name: str, arguments: str) -> st return f"call_{digest}" +_synthesize_oci_tool_call_id = synthesize_oci_tool_call_id + + def adapt_tools_to_openai_standard( tools: list[OCIToolCall], ) -> list[ChatCompletionMessageToolCall]: """Convert OCI tool-call objects in a response to the OpenAI format.""" return [ ChatCompletionMessageToolCall( - id=tool.id or _synthesize_oci_tool_call_id(i, tool.name, tool.arguments), + id=tool.id or synthesize_oci_tool_call_id(i, tool.name, tool.arguments), type="function", function={"name": tool.name, "arguments": tool.arguments}, ) @@ -341,7 +347,7 @@ def handle_generic_response( if response_message.toolCalls: message.tool_calls = adapt_tools_to_openai_standard(response_message.toolCalls) - model_response.choices[0].finish_reason = _normalize_oci_finish_reason(response_choice.finishReason) + model_response.choices[0].finish_reason = normalize_oci_finish_reason(response_choice.finishReason) oci_usage: Final = completion_response.chatResponse.usage reasoning_tokens: int | None = None @@ -408,7 +414,7 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: if typed_chunk.message and typed_chunk.message.toolCalls: tool_calls = [ { - "id": tc.id or _synthesize_oci_tool_call_id(i, tc.name, tc.arguments), + "id": tc.id or synthesize_oci_tool_call_id(i, tc.name, tc.arguments), "type": "function", "function": { "name": tc.name, @@ -418,7 +424,7 @@ def handle_generic_stream_chunk(dict_chunk: dict) -> ModelResponseStream: for i, tc in enumerate(typed_chunk.message.toolCalls) ] - finish_reason: Final[str | None] = _normalize_oci_finish_reason(typed_chunk.finishReason) + finish_reason: Final[str | None] = normalize_oci_finish_reason(typed_chunk.finishReason) return ModelResponseStream( choices=[ diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 7ec0e1c9cf6..14ef6225666 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -23,14 +23,14 @@ from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMExcepti from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, version, ) from litellm.llms.oci.chat.cohere import ( - _extract_text_content, adapt_messages_to_cohere_standard, adapt_tool_definitions_to_cohere_standard, + extract_text_content, handle_cohere_response, handle_cohere_stream_chunk, ) @@ -563,13 +563,13 @@ class OCIChatConfig(BaseConfig): system_messages: Final = [m for m in messages if m.get("role") == "system"] preamble_override = None if system_messages: - preamble: Final = "\n".join(_extract_text_content(m["content"]) for m in system_messages) + preamble: Final = "\n".join(extract_text_content(m["content"]) for m in system_messages) if preamble: preamble_override = preamble chat_request: Final = CohereChatRequest( apiFormat="COHERE", - message=_extract_text_content(user_messages[-1]["content"]), + message=extract_text_content(user_messages[-1]["content"]), chatHistory=adapt_messages_to_cohere_standard([m for m in messages if m.get("role") != "system"]), preambleOverride=preamble_override, **self._get_optional_params(OCIVendors.COHERE, optional_params, model), @@ -647,7 +647,7 @@ class OCIChatConfig(BaseConfig): timeout: float | httpx.Timeout | None = None, ) -> "OCIStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): - client = _get_httpx_client(params={}) + client = get_httpx_client(params={}) try: response: Final = client.post( diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index 9f46cbc5cd5..409e763dde1 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -31,7 +31,7 @@ def _reencode_as_jpeg(raw_image: bytes, original: str) -> str: return base64.b64encode(jpeg_image.getvalue()).decode("utf-8") -def _convert_image(image: str) -> str: +def convert_image(image: str) -> str: payload: Final = image.split(",")[-1] if image.startswith("data:") else image try: raw_image: Final = base64.b64decode(payload) @@ -42,6 +42,8 @@ def _convert_image(image: str) -> str: return _reencode_as_jpeg(raw_image, original=image) +_convert_image = convert_image + from litellm.llms.base_llm.base_utils import BaseLLMModelInfo diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index e307b8ca07a..c508c27c11a 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -41,7 +41,7 @@ from litellm.types.utils import ( StreamingChoices, ) -from ..common_utils import OllamaError, OllamaModelInfo, _convert_image +from ..common_utils import OllamaError, OllamaModelInfo, convert_image if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -465,7 +465,7 @@ class OllamaConfig(BaseConfig): if format is not None: data["format"] = format if images is not None: - data["images"] = [_convert_image(convert_to_ollama_image(image)) for image in images] + data["images"] = [convert_image(convert_to_ollama_image(image)) for image in images] if think is not None: data["think"] = think diff --git a/litellm/llms/oobabooga/chat/oobabooga.py b/litellm/llms/oobabooga/chat/oobabooga.py index cd118a0af29..e21a165d757 100644 --- a/litellm/llms/oobabooga/chat/oobabooga.py +++ b/litellm/llms/oobabooga/chat/oobabooga.py @@ -3,7 +3,7 @@ from collections.abc import Callable from typing import TYPE_CHECKING, Final import litellm -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client from litellm.utils import EmbeddingResponse, ModelResponse, Usage from ..common_utils import OobaboogaError @@ -65,7 +65,7 @@ def completion( additional_args={"complete_input_dict": data}, ) ## COMPLETION CALL - client: Final = _get_httpx_client() + client: Final = get_httpx_client() response: Final = client.post( completion_url, headers=headers, diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index 97e412942aa..94c5c5d016d 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -48,7 +48,7 @@ def _normalize_reasoning_effort_for_chat_completion( return None -def _get_effort_level(value: str | dict | None) -> str | None: +def get_effort_level(value: str | dict | None) -> str | None: """Extract the effective effort level from reasoning_effort (string or dict). Use this for guards that compare effort level (e.g. xhigh validation, "none" checks). @@ -64,6 +64,8 @@ def _get_effort_level(value: str | dict | None) -> str | None: return None +_get_effort_level = get_effort_level + GPT_REASONING_SERIES_MARKERS: Final = ("gpt-5", "gpt-6") @@ -163,6 +165,14 @@ class OpenAIGPT5Config(OpenAIGPTConfig): key=f"supports_{level}_reasoning_effort", ) + @classmethod + def supports_reasoning_effort_level( + cls, + model: str, + level: str, + ) -> bool: + return cls._supports_reasoning_effort_level(model, level) + @classmethod def effort_resolves_to_none(cls, model: str, effective_effort: str | None) -> bool: """Whether this request's reasoning effort ends up as "none", which is the single @@ -274,7 +284,7 @@ class OpenAIGPT5Config(OpenAIGPTConfig): # tool/sampling guards — dict inputs like {"effort": "none", "summary": "detailed"} # must be treated as effort="none" to avoid incorrect tool-drop or sampling errors. raw_reasoning_effort = non_default_params.get("reasoning_effort") or optional_params.get("reasoning_effort") - effective_effort: Final = _get_effort_level(raw_reasoning_effort) + effective_effort: Final = get_effort_level(raw_reasoning_effort) # Normalize dict reasoning_effort to string for Chat Completions API. # Example: {"effort": "high", "summary": "detailed"} -> "high" diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 6380c9390c6..936ec67ae92 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -387,6 +387,17 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): ) return hoisted_messages + def transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: mirrors override contract + model: str, + is_async: bool = False, + ) -> ( + list[AllMessageValues] # mutable-ok: mirrors override contract + | Coroutine[object, object, list[AllMessageValues]] + ): + return self._transform_messages(messages, model, is_async) + def remove_cache_control_flag_from_messages_and_tools( self, model: str, # allows overrides to selectively run this diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index b47edee9976..ffc5a5a71e0 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -315,7 +315,7 @@ class BaseOpenAILLM: # Get unified SSL configuration ssl_config: Final = get_ssl_configuration() - transport: Final = AsyncHTTPHandler._create_async_transport( + transport: Final = AsyncHTTPHandler.create_async_transport( ssl_context=(ssl_config if isinstance(ssl_config, ssl.SSLContext) else None), ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, shared_session=shared_session, @@ -324,11 +324,18 @@ class BaseOpenAILLM: return httpx.AsyncClient( verify=ssl_config, transport=transport, - mounts=AsyncHTTPHandler._create_httpx_proxy_mounts(transport, verify=ssl_config, cert=None), + mounts=AsyncHTTPHandler.create_httpx_proxy_mounts(transport, verify=ssl_config, cert=None), follow_redirects=True, http2=http2_enabled(), ) + @classmethod + def get_async_http_client( + cls, + shared_session: Optional["ClientSession"] = None, + ) -> httpx.AsyncClient | None: + return cls._get_async_http_client(shared_session) + @staticmethod def _get_sync_http_client() -> httpx.Client | None: if litellm.client_session is not None: diff --git a/litellm/llms/openai/completion/handler.py b/litellm/llms/openai/completion/handler.py index abacb0b815a..1bc752b6714 100644 --- a/litellm/llms/openai/completion/handler.py +++ b/litellm/llms/openai/completion/handler.py @@ -181,7 +181,7 @@ class OpenAITextCompletion(BaseLLM): openai_aclient = AsyncOpenAI( api_key=api_key, base_url=api_base, - http_client=BaseOpenAILLM._get_async_http_client(), + http_client=BaseOpenAILLM.get_async_http_client(), timeout=timeout, max_retries=max_retries, organization=organization, diff --git a/litellm/llms/openai/completion/transformation.py b/litellm/llms/openai/completion/transformation.py index a690a34f326..f7d2ce68082 100644 --- a/litellm/llms/openai/completion/transformation.py +++ b/litellm/llms/openai/completion/transformation.py @@ -11,7 +11,7 @@ from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUser from litellm.types.utils import Choices, Message, ModelResponse, TextCompletionResponse from ..chat.gpt_transformation import OpenAIGPTConfig -from .utils import _transform_prompt +from .utils import transform_prompt class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig): @@ -154,7 +154,7 @@ class OpenAITextCompletionConfig(BaseTextCompletionConfig, OpenAIGPTConfig): optional_params: dict, headers: dict, ) -> dict: - prompt: Final = _transform_prompt(messages) + prompt: Final = transform_prompt(messages) return { "model": model, "prompt": prompt, diff --git a/litellm/llms/openai/completion/utils.py b/litellm/llms/openai/completion/utils.py index 04892abe3cf..59ce066d2c2 100644 --- a/litellm/llms/openai/completion/utils.py +++ b/litellm/llms/openai/completion/utils.py @@ -22,7 +22,7 @@ def is_tokens_or_list_of_tokens(value: list): return False -def _transform_prompt( +def transform_prompt( messages: list[AllMessageValues] | list[OpenAITextCompletionUserMessage], ) -> AllPromptValues: if len(messages) == 1: # base case @@ -43,3 +43,6 @@ def _transform_prompt( raise e openai_prompt = prompt_str_list return openai_prompt + + +_transform_prompt = transform_prompt diff --git a/litellm/llms/openai/cost_calculation.py b/litellm/llms/openai/cost_calculation.py index 8c6bfe9796b..36d7e4d95de 100644 --- a/litellm/llms/openai/cost_calculation.py +++ b/litellm/llms/openai/cost_calculation.py @@ -100,7 +100,7 @@ def _video_resolution_to_cost_field_suffix(resolution: str) -> str | None: return safe -def _video_output_cost_per_second( +def video_output_cost_per_second( model_info: Mapping[str, Any], video_resolution: str | None, ) -> float | None: @@ -125,6 +125,9 @@ def _video_output_cost_per_second( return None +_video_output_cost_per_second = video_output_cost_per_second + + def video_generation_cost( model: str, duration_seconds: float, @@ -162,7 +165,7 @@ def video_generation_cost( ) return video_cost_per_second * duration_seconds - output_cost_per_second: Final = _video_output_cost_per_second(model_info, video_resolution) + output_cost_per_second: Final = video_output_cost_per_second(model_info, video_resolution) if output_cost_per_second is not None: verbose_logger.debug( "For model=%s - output_cost_per_second: %s; duration: %s", model, output_cost_per_second, duration_seconds diff --git a/litellm/llms/openai/fine_tuning/handler.py b/litellm/llms/openai/fine_tuning/handler.py index 1ff5909a103..c54263f4c35 100644 --- a/litellm/llms/openai/fine_tuning/handler.py +++ b/litellm/llms/openai/fine_tuning/handler.py @@ -48,10 +48,13 @@ def _normalize_fine_tuning_job_dict(data: dict[str, object], is_azure: bool = Fa return normalized -def _litellm_fine_tuning_job_from_response(response: FineTuningJob, is_azure: bool = False) -> LiteLLMFineTuningJob: +def litellm_fine_tuning_job_from_response(response: FineTuningJob, is_azure: bool = False) -> LiteLLMFineTuningJob: return LiteLLMFineTuningJob(**_normalize_fine_tuning_job_dict(response.model_dump(), is_azure=is_azure)) +_litellm_fine_tuning_job_from_response = litellm_fine_tuning_job_from_response + + class OpenAIFineTuningAPI: """ OpenAI methods to support for batches @@ -99,7 +102,7 @@ class OpenAIFineTuningAPI: ) -> LiteLLMFineTuningJob: response: Final = await openai_client.fine_tuning.jobs.create(**create_fine_tuning_job_data) - return _litellm_fine_tuning_job_from_response(response) + return litellm_fine_tuning_job_from_response(response) def create_fine_tuning_job( self, @@ -139,7 +142,7 @@ class OpenAIFineTuningAPI: ) verbose_logger.debug("creating fine tuning job, args= %s", create_fine_tuning_job_data) response: Final = cast(OpenAI, openai_client).fine_tuning.jobs.create(**create_fine_tuning_job_data) - return _litellm_fine_tuning_job_from_response(response) + return litellm_fine_tuning_job_from_response(response) async def acancel_fine_tuning_job( self, @@ -147,7 +150,7 @@ class OpenAIFineTuningAPI: openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response: Final = await openai_client.fine_tuning.jobs.cancel(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response) + return litellm_fine_tuning_job_from_response(response) def cancel_fine_tuning_job( self, @@ -187,7 +190,7 @@ class OpenAIFineTuningAPI: ) verbose_logger.debug("canceling fine tuning job, args= %s", fine_tuning_job_id) response: Final = cast(OpenAI, openai_client).fine_tuning.jobs.cancel(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response) + return litellm_fine_tuning_job_from_response(response) async def alist_fine_tuning_jobs( self, @@ -246,7 +249,7 @@ class OpenAIFineTuningAPI: openai_client: AsyncOpenAI | AsyncAzureOpenAI, ) -> LiteLLMFineTuningJob: response: Final = await openai_client.fine_tuning.jobs.retrieve(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response) + return litellm_fine_tuning_job_from_response(response) def retrieve_fine_tuning_job( self, @@ -286,4 +289,4 @@ class OpenAIFineTuningAPI: ) verbose_logger.debug("retrieving fine tuning job, id= %s", fine_tuning_job_id) response: Final = cast(OpenAI, openai_client).fine_tuning.jobs.retrieve(fine_tuning_job_id=fine_tuning_job_id) - return _litellm_fine_tuning_job_from_response(response) + return litellm_fine_tuning_job_from_response(response) diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 2461fb4e1af..a782bf35ac8 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -214,6 +214,13 @@ class OpenAIConfig(BaseConfig): def _transform_messages(self, messages: list[AllMessageValues], model: str) -> list[AllMessageValues]: return messages + def transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: mirrors override contract + model: str, + ) -> list[AllMessageValues]: # mutable-ok: mirrors override contract + return self._transform_messages(messages, model) + def map_openai_params( self, non_default_params: dict, @@ -464,6 +471,32 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): ) return client + def get_openai_client( + self, + is_async: bool, + api_key: str | None = None, + api_base: str | None = None, + api_version: str | None = None, + timeout: float | httpx.Timeout = (_get_openai_client.__defaults__ or ())[3], + max_retries: int | None = DEFAULT_MAX_RETRIES, + organization: str | None = None, + client: OpenAI | AsyncOpenAI | None = None, + shared_session: Optional["ClientSession"] = None, + litellm_params: Mapping[str, object] | None = None, + ) -> OpenAI | AsyncOpenAI | None: + return self._get_openai_client( + is_async, + api_key, + api_base, + api_version, + timeout, + max_retries, + organization, + client, + shared_session, + litellm_params, + ) + @track_llm_api_timing() async def make_openai_chat_completion_request( self, @@ -799,7 +832,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): api_key=openai_client.api_key, additional_args={ "headers": headers, - "api_base": openai_client._base_url._uri_reference, + "api_base": openai_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "acompletion": acompletion, "complete_input_dict": data, "openai_sdk": True, @@ -942,7 +975,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): api_key=openai_aclient.api_key, additional_args={ "headers": {"Authorization": f"Bearer {openai_aclient.api_key}"}, - "api_base": openai_aclient._base_url._uri_reference, + "api_base": openai_aclient._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "acompletion": True, "complete_input_dict": data, "openai_sdk": True, @@ -1052,7 +1085,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): api_key=api_key, additional_args={ "headers": {"Authorization": f"Bearer {openai_client.api_key}"}, - "api_base": openai_client._base_url._uri_reference, + "api_base": openai_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "acompletion": False, "complete_input_dict": data, }, @@ -1520,7 +1553,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): api_key=openai_client.api_key, additional_args={ "headers": {"Authorization": f"Bearer {openai_client.api_key}"}, - "api_base": openai_client._base_url._uri_reference, + "api_base": openai_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "acompletion": True, "complete_input_dict": data, }, diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index 182b24c9901..325169cfe15 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -108,6 +108,13 @@ class OpenAIRealtime(OpenAIChatCompletion): url = url.copy_with(params=upstream_params) return str(url) + def construct_url( + self, + api_base: str, + query_params: RealtimeQueryParams, + ) -> str: + return self._construct_url(api_base, query_params) + def _make_event_normalizer(self) -> RealtimeEventNormalizer | None: """Return a per-session GA event normalizer, or None for passthrough. diff --git a/litellm/llms/openai/transcriptions/handler.py b/litellm/llms/openai/transcriptions/handler.py index 014251db821..733d7930144 100644 --- a/litellm/llms/openai/transcriptions/handler.py +++ b/litellm/llms/openai/transcriptions/handler.py @@ -113,7 +113,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): input=None, api_key=openai_client.api_key, additional_args={ - "api_base": openai_client._base_url._uri_reference, + "api_base": openai_client._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "atranscription": True, "complete_input_dict": data, }, @@ -176,7 +176,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion): input=None, api_key=openai_aclient.api_key, additional_args={ - "api_base": openai_aclient._base_url._uri_reference, + "api_base": openai_aclient._base_url._uri_reference, # pyright: ignore[reportPrivateUsage] # SDK URL internals "atranscription": True, "complete_input_dict": data, }, diff --git a/litellm/llms/openai_like/chat/handler.py b/litellm/llms/openai_like/chat/handler.py index 855c49c320b..62a84ad3f15 100644 --- a/litellm/llms/openai_like/chat/handler.py +++ b/litellm/llms/openai_like/chat/handler.py @@ -268,7 +268,7 @@ class OpenAILikeChatHandler(OpenAILikeBase): model=model, provider=LlmProviders(custom_llm_provider) ) if isinstance(provider_config, OpenAIGPTConfig) or isinstance(provider_config, OpenAIConfig): - messages = provider_config._transform_messages(messages=messages, model=model) + messages = provider_config.transform_messages(messages=messages, model=model) data: Final = { "model": model, diff --git a/litellm/llms/openai_like/chat/transformation.py b/litellm/llms/openai_like/chat/transformation.py index b0770f1a7e1..83fcb43c9b8 100644 --- a/litellm/llms/openai_like/chat/transformation.py +++ b/litellm/llms/openai_like/chat/transformation.py @@ -31,6 +31,13 @@ class OpenAILikeChatConfig(OpenAIGPTConfig): dynamic_api_key = api_key or get_secret_str("OPENAI_LIKE_API_KEY") or "" # vllm does not require an api key return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + @staticmethod def _json_mode_convert_tool_response_to_message( message: ChatCompletionAssistantMessage, json_mode: bool diff --git a/litellm/llms/openai_like/dynamic_config.py b/litellm/llms/openai_like/dynamic_config.py index b8517a4257c..f0bf710b951 100644 --- a/litellm/llms/openai_like/dynamic_config.py +++ b/litellm/llms/openai_like/dynamic_config.py @@ -71,6 +71,13 @@ def create_config_class(provider: SimpleProviderConfig): return resolved_base, resolved_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_complete_url( self, api_base: str | None, diff --git a/litellm/llms/perplexity/chat/transformation.py b/litellm/llms/perplexity/chat/transformation.py index dca2f9857b8..2894218da67 100644 --- a/litellm/llms/perplexity/chat/transformation.py +++ b/litellm/llms/perplexity/chat/transformation.py @@ -30,6 +30,13 @@ class PerplexityChatConfig(OpenAIGPTConfig): dynamic_api_key = api_key or get_secret_str("PERPLEXITYAI_API_KEY") or get_secret_str("PERPLEXITY_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_supported_openai_params(self, model: str) -> list: """ Perplexity supports a subset of OpenAI params diff --git a/litellm/llms/petals/completion/handler.py b/litellm/llms/petals/completion/handler.py index c7cfeb1dd1a..19010a6722a 100644 --- a/litellm/llms/petals/completion/handler.py +++ b/litellm/llms/petals/completion/handler.py @@ -10,7 +10,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, + get_httpx_client, ) from litellm.utils import ModelResponse, Usage @@ -66,7 +66,7 @@ def completion( ## COMPLETION CALL if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client() + client = get_httpx_client() response: Final = client.post(api_base, data=data) ## LOGGING diff --git a/litellm/llms/ragflow/chat/transformation.py b/litellm/llms/ragflow/chat/transformation.py index 964428fdb1a..97b723f7dd5 100644 --- a/litellm/llms/ragflow/chat/transformation.py +++ b/litellm/llms/ragflow/chat/transformation.py @@ -157,6 +157,15 @@ class RAGFlowConfig(OpenAIConfig): return dynamic_api_base, dynamic_api_key, custom_llm_provider + def get_openai_compatible_provider_info( + self, + model: str, + api_base: str | None, + api_key: str | None, + custom_llm_provider: str, + ) -> tuple[str | None, str | None, str]: + return self._get_openai_compatible_provider_info(model, api_base, api_key, custom_llm_provider) + def validate_environment( self, headers: dict, diff --git a/litellm/llms/replicate/chat/handler.py b/litellm/llms/replicate/chat/handler.py index 764ea40db8d..e4233ca725d 100644 --- a/litellm/llms/replicate/chat/handler.py +++ b/litellm/llms/replicate/chat/handler.py @@ -12,8 +12,8 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.openai import AllMessageValues from litellm.utils import CustomStreamWrapper, ModelResponse @@ -186,7 +186,7 @@ def completion( ) ## COMPLETION CALL - httpx_client: Final = _get_httpx_client( + httpx_client: Final = get_httpx_client( params={"timeout": 600.0}, ) response = httpx_client.post( diff --git a/litellm/llms/runwayml/image_generation/transformation.py b/litellm/llms/runwayml/image_generation/transformation.py index e5e988328d8..ddc1b9474ac 100644 --- a/litellm/llms/runwayml/image_generation/transformation.py +++ b/litellm/llms/runwayml/image_generation/transformation.py @@ -220,9 +220,9 @@ class RunwayMLImageGenerationConfig(BaseImageGenerationConfig): Returns: Final response with completed task """ - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client - client: Final = _get_httpx_client() + client: Final = get_httpx_client() start_time: Final = time.time() # Build task status URL diff --git a/litellm/llms/runwayml/text_to_speech/transformation.py b/litellm/llms/runwayml/text_to_speech/transformation.py index 6769accc1d6..d41ffea17dc 100644 --- a/litellm/llms/runwayml/text_to_speech/transformation.py +++ b/litellm/llms/runwayml/text_to_speech/transformation.py @@ -305,9 +305,9 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): Returns: Final response with completed task """ - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client - client: Final = _get_httpx_client() + client: Final = get_httpx_client() start_time: Final = time.time() # Build task status URL @@ -506,9 +506,9 @@ class RunwayMLTextToSpeechConfig(BaseTextToSpeechConfig): raise ValueError(f"RunwayML TTS audio URL is not a string: {audio_url}") # Download the audio file - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client - client: Final = _get_httpx_client() + client: Final = get_httpx_client() audio_response: Final = client.get(url=audio_url) audio_response.raise_for_status() diff --git a/litellm/llms/runwayml/videos/transformation.py b/litellm/llms/runwayml/videos/transformation.py index 1abdece44f2..ec7f0102bd6 100644 --- a/litellm/llms/runwayml/videos/transformation.py +++ b/litellm/llms/runwayml/videos/transformation.py @@ -15,8 +15,8 @@ from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.secret_managers.main import get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -464,7 +464,7 @@ class RunwayMLVideoConfig(BaseVideoConfig): video_url: Final = self._extract_video_url_from_response(response_data) # Download the video from the CloudFront URL synchronously - httpx_client: Final[HTTPHandler] = _get_httpx_client() + httpx_client: Final[HTTPHandler] = get_httpx_client() video_response: Final = httpx_client.get(video_url) video_response.raise_for_status() diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index 8f320343de1..6dfc274fed1 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -21,8 +21,8 @@ from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import LlmProviders @@ -155,7 +155,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: if client is None or isinstance(client, AsyncHTTPHandler): - client = _get_httpx_client(params={}) + client = get_httpx_client(params={}) try: response: Final = client.post( diff --git a/litellm/llms/sagemaker/completion/handler.py b/litellm/llms/sagemaker/completion/handler.py index 3e110a869bc..46d204cf385 100644 --- a/litellm/llms/sagemaker/completion/handler.py +++ b/litellm/llms/sagemaker/completion/handler.py @@ -12,8 +12,8 @@ from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, pop_aws_auth_params from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.openai import AllMessageValues from litellm.utils import ( @@ -253,7 +253,7 @@ class SagemakerLLM(BaseAWSLLM): ## LOGGING timeout = 300.0 - sync_handler: Final = _get_httpx_client() + sync_handler: Final = get_httpx_client() ## LOGGING logging_obj.pre_call( input=[], @@ -317,7 +317,7 @@ class SagemakerLLM(BaseAWSLLM): client=None, ): if client is None: - client = _get_httpx_client() + client = get_httpx_client() sync_response: Final = client.post( api_base, headers=headers, diff --git a/litellm/llms/sagemaker/embedding/cohere_transformation.py b/litellm/llms/sagemaker/embedding/cohere_transformation.py index 4687ff6b3f4..35352f65dda 100644 --- a/litellm/llms/sagemaker/embedding/cohere_transformation.py +++ b/litellm/llms/sagemaker/embedding/cohere_transformation.py @@ -79,7 +79,7 @@ class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): input_list = [str(input)] return dict( - BedrockCohereEmbeddingConfig()._transform_request( + BedrockCohereEmbeddingConfig().transform_request( model=model, input=input_list, inference_params=optional_params, @@ -111,7 +111,7 @@ class SagemakerCohereEmbeddingConfig(BaseEmbeddingConfig): if isinstance(input_value, str): input_value = [input_value] - return CohereEmbeddingConfig()._populate_embedding_response( + return CohereEmbeddingConfig().populate_embedding_response( response_json=raw_response.json(), model_response=model_response, model=model, diff --git a/litellm/llms/sap/credentials.py b/litellm/llms/sap/credentials.py index b5ff34b81d7..452dcc15c4c 100644 --- a/litellm/llms/sap/credentials.py +++ b/litellm/llms/sap/credentials.py @@ -15,7 +15,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict import litellm from litellm._logging import verbose_logger -from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client +from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_httpx_client AUTH_ENDPOINT_SUFFIX: Final = "/oauth/token" @@ -388,7 +388,7 @@ def _request_token( handler = HTTPHandler(client=raw_client) resp = handler.post(auth_url, data=data, timeout=timeout) return _bearer_token_and_expiry(resp) - handler = _get_httpx_client() + handler = get_httpx_client() resp = handler.post(auth_url, data=data, timeout=timeout) return _bearer_token_and_expiry(resp) except Exception as e: diff --git a/litellm/llms/snowflake/utils.py b/litellm/llms/snowflake/utils.py index 26cbd246346..a8c14d1d7f3 100644 --- a/litellm/llms/snowflake/utils.py +++ b/litellm/llms/snowflake/utils.py @@ -118,3 +118,10 @@ class SnowflakeBaseConfig: ) -> tuple[str | None, str | None]: dynamic_api_key: Final = api_key or get_secret_str("SNOWFLAKE_JWT") return api_base, dynamic_api_key + + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) diff --git a/litellm/llms/soniox/audio_transcription/handler.py b/litellm/llms/soniox/audio_transcription/handler.py index a9723125f17..b53955ae75d 100644 --- a/litellm/llms/soniox/audio_transcription/handler.py +++ b/litellm/llms/soniox/audio_transcription/handler.py @@ -31,8 +31,8 @@ from litellm.litellm_core_utils.audio_utils.utils import ( from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.llms.soniox.audio_transcription.transformation import ( SonioxAudioTranscriptionConfig, @@ -384,7 +384,7 @@ class SonioxAudioTranscriptionHandler: client if isinstance(client, HTTPHandler) else ( - _get_httpx_client( + get_httpx_client( params={"ssl_verify": litellm_params.get("ssl_verify", None)}, ) ) @@ -449,7 +449,7 @@ class SonioxAudioTranscriptionHandler: fetched: Final[_SonioxJsonView] = {"transcript": transcript_resp.json()} payload: Final = {"transcription": transcription_meta, "transcript": fetched["transcript"]} - response: Final = provider_config._build_response_from_payload( + response: Final = provider_config.build_response_from_payload( payload, model_response=model_response, response_format=handler_opts.get("response_format"), @@ -683,7 +683,7 @@ class SonioxAudioTranscriptionHandler: fetched: Final[_SonioxJsonView] = {"transcript": transcript_resp.json()} payload: Final = {"transcription": transcription_meta, "transcript": fetched["transcript"]} - response: Final = provider_config._build_response_from_payload( + response: Final = provider_config.build_response_from_payload( payload, model_response=model_response, response_format=handler_opts.get("response_format"), diff --git a/litellm/llms/soniox/audio_transcription/transformation.py b/litellm/llms/soniox/audio_transcription/transformation.py index 48510015f44..aa1ff239320 100644 --- a/litellm/llms/soniox/audio_transcription/transformation.py +++ b/litellm/llms/soniox/audio_transcription/transformation.py @@ -269,3 +269,11 @@ class SonioxAudioTranscriptionConfig(BaseAudioTranscriptionConfig): } ) return response + + def build_response_from_payload( + self, + payload: dict[str, object], # mutable-ok: mirrors override contract + model_response: TranscriptionResponse | None = None, + response_format: str | None = None, + ) -> TranscriptionResponse: + return self._build_response_from_payload(payload, model_response, response_format) diff --git a/litellm/llms/together_ai/completion/transformation.py b/litellm/llms/together_ai/completion/transformation.py index c58e29ad88d..67e2cc8c349 100644 --- a/litellm/llms/together_ai/completion/transformation.py +++ b/litellm/llms/together_ai/completion/transformation.py @@ -16,7 +16,7 @@ from litellm.types.llms.openai import ( ) from ...openai.completion.transformation import OpenAITextCompletionConfig -from ...openai.completion.utils import _transform_prompt +from ...openai.completion.utils import transform_prompt class TogetherAITextCompletionConfig(OpenAITextCompletionConfig): @@ -27,7 +27,7 @@ class TogetherAITextCompletionConfig(OpenAITextCompletionConfig): """ TogetherAI expects a string prompt. """ - initial_prompt: Final[AllPromptValues] = _transform_prompt(messages) + initial_prompt: Final[AllPromptValues] = transform_prompt(messages) ## TOGETHER AI SPECIFIC VALIDATION ## if isinstance(initial_prompt, list) and is_tokens_or_list_of_tokens(value=initial_prompt): raise ValueError("TogetherAI does not support integers as input") diff --git a/litellm/llms/together_ai/rerank/handler.py b/litellm/llms/together_ai/rerank/handler.py index 2fff0854f10..ebf6c77a9ef 100644 --- a/litellm/llms/together_ai/rerank/handler.py +++ b/litellm/llms/together_ai/rerank/handler.py @@ -11,8 +11,8 @@ from pydantic import ConfigDict, TypeAdapter import litellm from litellm.llms.base import BaseLLM from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.llms.together_ai.rerank.transformation import TogetherAIRerankConfig from litellm.types.rerank import RerankRequest, RerankResponse @@ -38,7 +38,7 @@ class TogetherAIRerank(BaseLLM): max_chunks_per_doc: int | None = None, _is_async: bool | None = False, ) -> RerankResponse: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() request_data: Final = RerankRequest( model=model, @@ -72,7 +72,7 @@ class TogetherAIRerank(BaseLLM): _json_response: Final = _JSON_DICT.validate_python(response.json()) - return TogetherAIRerankConfig()._transform_response(_json_response) + return TogetherAIRerankConfig().transform_response(_json_response) async def async_rerank( # New async method self, @@ -97,4 +97,4 @@ class TogetherAIRerank(BaseLLM): _json_response: Final = _JSON_DICT.validate_python(response.json()) - return TogetherAIRerankConfig()._transform_response(_json_response) + return TogetherAIRerankConfig().transform_response(_json_response) diff --git a/litellm/llms/together_ai/rerank/transformation.py b/litellm/llms/together_ai/rerank/transformation.py index 29551eaccd3..f1c3ca9e87d 100644 --- a/litellm/llms/together_ai/rerank/transformation.py +++ b/litellm/llms/together_ai/rerank/transformation.py @@ -56,3 +56,9 @@ class TogetherAIRerankConfig: results=rerank_results, meta=rerank_meta, ) # Return response + + def transform_response( + self, + response: dict[str, object], # mutable-ok: mirrors override contract + ) -> RerankResponse: + return self._transform_response(response) diff --git a/litellm/llms/v0/chat/transformation.py b/litellm/llms/v0/chat/transformation.py index 28f8c6cf342..81ee06a7501 100644 --- a/litellm/llms/v0/chat/transformation.py +++ b/litellm/llms/v0/chat/transformation.py @@ -28,6 +28,13 @@ class V0ChatConfig(OpenAILikeChatConfig): dynamic_api_key: Final = api_key or get_secret_str("V0_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def get_supported_openai_params(self, model: str) -> list: """ v0 supports a limited subset of OpenAI parameters diff --git a/litellm/llms/vercel_ai_gateway/chat/transformation.py b/litellm/llms/vercel_ai_gateway/chat/transformation.py index a6fd7e5a7b5..19c4c2659ca 100644 --- a/litellm/llms/vercel_ai_gateway/chat/transformation.py +++ b/litellm/llms/vercel_ai_gateway/chat/transformation.py @@ -37,6 +37,13 @@ class VercelAIGatewayConfig(OpenAIGPTConfig): user_api_key = api_key or get_secret_str("VERCEL_AI_GATEWAY_API_KEY") or get_secret_str("VERCEL_OIDC_TOKEN") return api_base, user_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def map_openai_params( self, non_default_params: dict, diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py index 81b2c3d5109..10e89165b73 100644 --- a/litellm/llms/vertex_ai/agent_engine/transformation.py +++ b/litellm/llms/vertex_ai/agent_engine/transformation.py @@ -373,12 +373,12 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): """Get a CustomStreamWrapper for synchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, - _get_httpx_client, + get_httpx_client, ) from litellm.utils import CustomStreamWrapper if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={}) + client = get_httpx_client(params={}) # Avoid logging sensitive api_base directly verbose_logger.debug("Making sync streaming request to Vertex AI endpoint.") diff --git a/litellm/llms/vertex_ai/batches/handler.py b/litellm/llms/vertex_ai/batches/handler.py index a15ea4d845b..de5fe39cfc9 100644 --- a/litellm/llms/vertex_ai/batches/handler.py +++ b/litellm/llms/vertex_ai/batches/handler.py @@ -14,12 +14,12 @@ from litellm.litellm_core_utils.url_utils import ( ) from litellm.llms.custom_httpx.http_handler import ( HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.llms.vertex_ai.common_utils import VertexAIError, get_vertex_base_url from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM -from litellm.llms.vertex_ai.vertex_llm_base import _graft_default_vertex_path +from litellm.llms.vertex_ai.vertex_llm_base import graft_default_vertex_path from litellm.types.llms.openai import CreateBatchRequest from litellm.types.llms.vertex_ai import ( VERTEX_CREDENTIALS_TYPES, @@ -106,7 +106,7 @@ class VertexAIBatchPrediction(VertexLLM): "use a publisher model or fine-tuned Gemini endpoint deployment instead." ), ) - sync_handler: Final = _get_httpx_client() + sync_handler: Final = get_httpx_client() access_token, project_id = self._ensure_access_token( credentials=vertex_credentials, @@ -193,7 +193,7 @@ class VertexAIBatchPrediction(VertexLLM): return default_endpoint_url api_base_path: Final = urlparse(api_base).path.rstrip("/") if api_base_path in ("/v1", "/v1beta1"): - return _graft_default_vertex_path(api_base=api_base, default_url=default_endpoint_url) + return graft_default_vertex_path(api_base=api_base, default_url=default_endpoint_url) return api_base.rstrip("/") + urlparse(default_endpoint_url).path def _resolve_fine_tuned_endpoint_model( @@ -302,7 +302,7 @@ class VertexAIBatchPrediction(VertexLLM): max_retries: int | None, logging_obj: "LiteLLMLoggingObj | None" = None, ) -> LiteLLMBatch | Coroutine[object, object, LiteLLMBatch]: - sync_handler: Final = _get_httpx_client() + sync_handler: Final = get_httpx_client() access_token, project_id = self._ensure_access_token( credentials=vertex_credentials, @@ -462,7 +462,7 @@ class VertexAIBatchPrediction(VertexLLM): timeout: float | httpx.Timeout, max_retries: int | None, ): - sync_handler: Final = _get_httpx_client() + sync_handler: Final = get_httpx_client() access_token, project_id = self._ensure_access_token( credentials=vertex_credentials, @@ -610,7 +610,7 @@ class VertexAIBatchPrediction(VertexLLM): timeout=timeout, ) - sync_handler: Final = _get_httpx_client() + sync_handler: Final = get_httpx_client() try: sync_handler.post( url=api_base, diff --git a/litellm/llms/vertex_ai/batches/transformation.py b/litellm/llms/vertex_ai/batches/transformation.py index 11c99d130be..52ea3244cf2 100644 --- a/litellm/llms/vertex_ai/batches/transformation.py +++ b/litellm/llms/vertex_ai/batches/transformation.py @@ -9,7 +9,7 @@ from litellm._logging import verbose_logger from litellm._uuid import uuid from litellm.llms.vertex_ai.common_utils import ( VertexAIError, - _convert_vertex_datetime_to_openai_datetime, + convert_vertex_datetime_to_openai_datetime, ) from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest from litellm.types.llms.vertex_ai import * @@ -181,9 +181,9 @@ class VertexAIBatchTransformation: cls, response: VertexBatchPredictionResponse ) -> LiteLLMBatch: return LiteLLMBatch( - id=cls._get_batch_id_from_vertex_ai_batch_response(response), + id=cls.get_batch_id_from_vertex_ai_batch_response(response), completion_window="24h", - created_at=_convert_vertex_datetime_to_openai_datetime(vertex_datetime=response.get("createTime", "")), + created_at=convert_vertex_datetime_to_openai_datetime(vertex_datetime=response.get("createTime", "")), endpoint="", input_file_id=cls._get_input_file_id_from_vertex_ai_batch_response(response), object="batch", @@ -217,7 +217,7 @@ class VertexAIBatchTransformation: } @classmethod - def _get_batch_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str: + def get_batch_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str: """ Gets the batch id from the Vertex AI Batch response safely @@ -232,6 +232,8 @@ class VertexAIBatchTransformation: parts: Final = _name.split("/") return parts[-1] if parts else _name + _get_batch_id_from_vertex_ai_batch_response = get_batch_id_from_vertex_ai_batch_response + @classmethod def _get_input_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str: """ diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 9a0209b87d9..8c77413bace 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -292,7 +292,7 @@ def get_supports_system_message( supports_system_message = supports_system_messages(model=model, custom_llm_provider=_custom_llm_provider) # Vertex Models called in the `/gemini` request/response format also support system messages - if litellm.VertexGeminiConfig._is_model_gemini_spec_model(model): + if litellm.VertexGeminiConfig.is_model_gemini_spec_model(model): supports_system_message = True except Exception as e: verbose_logger.warning( @@ -498,7 +498,7 @@ def _get_embedding_url( return url, endpoint -def _get_vertex_url( +def get_vertex_url( mode: all_gemini_url_modes, model: str, stream: bool | None, @@ -556,7 +556,10 @@ def _get_vertex_url( return url, endpoint -def _get_gemini_url( +_get_vertex_url = get_vertex_url + + +def get_gemini_url( mode: all_gemini_url_modes, model: str, stream: bool | None, @@ -572,7 +575,7 @@ def _get_gemini_url( ) _gemini_model_name: Final = f"models/{model}" - api_version: Final = "v1alpha" if VertexGeminiConfig._is_gemini_3_or_newer(model) else "v1beta" + api_version: Final = "v1alpha" if VertexGeminiConfig.is_gemini_3_or_newer(model) else "v1beta" if mode == "chat": endpoint = "generateContent" @@ -600,7 +603,10 @@ def _get_gemini_url( return url, endpoint -def _check_text_in_content(parts: list[PartType]) -> bool: +_get_gemini_url = get_gemini_url + + +def check_text_in_content(parts: list[PartType]) -> bool: """ check that user_content has 'text' parameter. - Known Vertex Error: Unable to submit request because it must have a text parameter. @@ -615,6 +621,9 @@ def _check_text_in_content(parts: list[PartType]) -> bool: return has_text_param +_check_text_in_content = check_text_in_content + + def _fix_enum_empty_strings(schema, depth=0): """Fix empty strings in enum values by replacing them with None. Gemini doesn't accept empty strings in enums.""" if depth > DEFAULT_MAX_RECURSE_DEPTH: @@ -683,7 +692,7 @@ def _fix_enum_types(schema, depth=0): _fix_enum_types(item, depth=depth + 1) -def _build_vertex_schema(parameters: dict, add_property_ordering: bool = False): +def build_vertex_schema(parameters: dict, add_property_ordering: bool = False): """ This is a modified version of https://github.com/google-gemini/generative-ai-python/blob/8f77cc6ac99937cd3a81299ecf79608b91b06bbb/google/generativeai/types/content_types.py#L419 @@ -734,7 +743,10 @@ def _build_vertex_schema(parameters: dict, add_property_ordering: bool = False): return parameters -def _build_json_schema(parameters: dict) -> dict: +_build_vertex_schema = build_vertex_schema + + +def build_json_schema(parameters: dict) -> dict: """ Build a JSON Schema for use with Gemini's responseJsonSchema parameter. @@ -760,6 +772,9 @@ def _build_json_schema(parameters: dict) -> dict: return parameters +_build_json_schema = build_json_schema + + def _filter_anyof_fields(schema_dict: dict[str, object]) -> dict[str, object]: """ When anyof is present, only keep the anyof field and its contents - otherwise VertexAI will throw an error - https://github.com/BerriAI/litellm/issues/11164 @@ -969,7 +984,7 @@ def strip_field(schema, field_name: str): strip_field(items, field_name) -def _convert_vertex_datetime_to_openai_datetime(vertex_datetime: str) -> int: +def convert_vertex_datetime_to_openai_datetime(vertex_datetime: str) -> int: """ Converts a Vertex AI datetime string to an OpenAI datetime integer @@ -984,6 +999,9 @@ def _convert_vertex_datetime_to_openai_datetime(vertex_datetime: str) -> int: return int(dt.timestamp()) +_convert_vertex_datetime_to_openai_datetime = convert_vertex_datetime_to_openai_datetime + + def _convert_schema_types(schema, depth=0): """ Convert type arrays and lowercase types for Vertex AI compatibility. @@ -1310,11 +1328,11 @@ class VertexAITokenCounter(BaseTokenCounter): else: from litellm.llms.vertex_ai.count_tokens.handler import VertexAITokenCounter from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, # pyright: ignore[reportPrivateUsage] # shared helper already used by gemini/chat, context_caching, and vertex_and_google_ai_studio_gemini + gemini_convert_messages_with_history, # pyright: ignore[reportPrivateUsage] # shared helper already used by gemini/chat, context_caching, and vertex_and_google_ai_studio_gemini ) resolved_contents: Final = ( - contents if contents is not None else _gemini_convert_messages_with_history(messages=messages or []) + contents if contents is not None else gemini_convert_messages_with_history(messages=messages or []) ) count_tokens_params: Final = { diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index d5478920de0..f683ad0e377 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -16,8 +16,8 @@ from litellm.utils import is_cached_message from ..common_utils import get_supports_system_message from ..gemini.transformation import ( - _gemini_convert_messages_with_history, - _transform_system_message, + gemini_convert_messages_with_history, + transform_system_message, ) @@ -171,11 +171,11 @@ def transform_openai_messages_to_gemini_context_caching( supports_system_message: Final = get_supports_system_message(model=model, custom_llm_provider=custom_llm_provider) - transformed_system_messages, new_messages = _transform_system_message( + transformed_system_messages, new_messages = transform_system_message( supports_system_message=supports_system_message, messages=messages ) - transformed_messages: Final = _gemini_convert_messages_with_history( + transformed_messages: Final = gemini_convert_messages_with_history( messages=new_messages, model=model, custom_llm_provider=custom_llm_provider, diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index a7e0bef37df..da5b974f186 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -45,10 +45,10 @@ from litellm.llms.base_llm.files.transformation import ( ) from litellm.llms.vertex_ai.batches.transformation import vertex_embedding_prompt_token_count from litellm.llms.vertex_ai.common_utils import ( - _convert_vertex_datetime_to_openai_datetime, + convert_vertex_datetime_to_openai_datetime, get_vertex_ai_fine_tuned_endpoint_id, ) -from litellm.llms.vertex_ai.gemini.transformation import _transform_request_body +from litellm.llms.vertex_ai.gemini.transformation import transform_request_body from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) @@ -681,7 +681,7 @@ def _openai_batch_jsonl_entry_to_vertex_rows( if _is_responses_batch_entry(openai_entry) else openai_request_body ) - vertex_request_body: Final = _transform_request_body( + vertex_request_body: Final = transform_request_body( messages=map_developer_role_to_system_role(chat_request_body.get("messages", [])), model=chat_request_body.get("model", ""), optional_params=map_openai_to_vertex_params(chat_request_body), @@ -1125,7 +1125,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): purpose=response_object.get("purpose", "batch"), id=f"gs://{gcs_id}", filename=response_object.get("name", ""), - created_at=_convert_vertex_datetime_to_openai_datetime( + created_at=convert_vertex_datetime_to_openai_datetime( vertex_datetime=response_object.get("timeCreated", "") ), status="uploaded", @@ -1172,9 +1172,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): return OpenAIFileObject( id=f"gs://{gcs_id}", bytes=int(response_json.get("size", 0)), - created_at=_convert_vertex_datetime_to_openai_datetime( - vertex_datetime=response_json.get("timeCreated", "") - ), + created_at=convert_vertex_datetime_to_openai_datetime(vertex_datetime=response_json.get("timeCreated", "")), filename=response_json.get("name", ""), object="file", purpose=response_json.get("metadata", {}).get("purpose", "batch"), @@ -1490,7 +1488,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): # Use existing VertexGeminiConfig transformation model_response: Final = ModelResponse() - transformed_response = vertex_gemini_config._transform_google_generate_content_to_openai_model_response( + transformed_response = vertex_gemini_config.transform_google_generate_content_to_openai_model_response( completion_response=vertex_response, model_response=model_response, model=model, diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 51b87233c30..8749b81d7f1 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -57,7 +57,7 @@ from litellm.types.llms.vertex_ai import ( from litellm.types.utils import GenericImageParsingChunk, LlmProviders from ..common_utils import ( - _check_text_in_content, + check_text_in_content, get_supports_response_schema, get_supports_system_message, ) @@ -193,7 +193,7 @@ def _apply_gemini_metadata( part_dict: Final = dict(part) - if media_resolution_enum is not None and VertexGeminiConfig._is_gemini_3_or_newer(model): + if media_resolution_enum is not None and VertexGeminiConfig.is_gemini_3_or_newer(model): part_dict["media_resolution"] = media_resolution_enum if video_metadata is not None: @@ -584,17 +584,23 @@ def _process_gemini_media( raise e -def _snake_to_camel(snake_str: str) -> str: +def snake_to_camel(snake_str: str) -> str: """Convert snake_case to camelCase""" components: Final = snake_str.split("_") return components[0] + "".join(x.capitalize() for x in components[1:]) -def _camel_to_snake(camel_str: str) -> str: +_snake_to_camel = snake_to_camel + + +def camel_to_snake(camel_str: str) -> str: """Convert camelCase to snake_case""" return re.sub(r"(? str | None: """ Get the equivalent key from available keys, checking both camelCase and snake_case variants @@ -603,12 +609,12 @@ def _get_equivalent_key(key: str, available_keys: set) -> str | None: return key # Try camelCase version - camel_key: Final = _snake_to_camel(key) + camel_key: Final = snake_to_camel(key) if camel_key in available_keys: return camel_key # Try snake_case version - snake_key: Final = _camel_to_snake(key) + snake_key: Final = camel_to_snake(key) if snake_key in available_keys: return snake_key @@ -687,7 +693,7 @@ def _collect_tool_call_thought_signatures( return frozenset(signatures) -def _gemini_convert_messages_with_history( +def gemini_convert_messages_with_history( messages: list[AllMessageValues], model: str | None = None, litellm_params: dict | None = None, @@ -715,7 +721,7 @@ def _gemini_convert_messages_with_history( from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig - forward_function_call_id: Final = VertexGeminiConfig._forward_gemini_function_call_id(model or "") + forward_function_call_id: Final = VertexGeminiConfig.forward_gemini_function_call_id(model or "") try: while msg_i < len(messages): @@ -861,7 +867,7 @@ def _gemini_convert_messages_with_history( - Known Vertex Error: Unable to submit request because it must have a text parameter. - Relevant Issue: https://github.com/BerriAI/litellm/issues/5515 """ - has_text_in_content = _check_text_in_content(user_content) + has_text_in_content = check_text_in_content(user_content) if has_text_in_content is False: verbose_logger.warning( "No text in user content. Adding a blank text to user content, to ensure Gemini doesn't fail the request. Relevant Issue - https://github.com/BerriAI/litellm/issues/5515" @@ -1053,6 +1059,8 @@ def _gemini_convert_messages_with_history( raise e +_gemini_convert_messages_with_history = gemini_convert_messages_with_history + # Keys that LiteLLM consumes internally and must never be forwarded to the _LITELLM_INTERNAL_EXTRA_BODY_KEYS: Final[frozenset] = frozenset({"cache", "tags"}) @@ -1124,7 +1132,7 @@ def _rewrite_google_maps_response_format(data: RequestBody) -> None: _rewrite_mime_type_to_response_format(generation_config) -def _transform_request_body( +def transform_request_body( messages: list[AllMessageValues], model: str, optional_params: dict, @@ -1137,7 +1145,7 @@ def _transform_request_body( """ # Separate system prompt from rest of message supports_system_message: Final = get_supports_system_message(model=model, custom_llm_provider=custom_llm_provider) - system_instructions, messages = _transform_system_message( + system_instructions, messages = transform_system_message( supports_system_message=supports_system_message, messages=messages ) # Checks for 'response_schema' support - if passed in @@ -1163,11 +1171,11 @@ def _transform_request_body( try: if custom_llm_provider == "gemini": - content = litellm.GoogleAIStudioGeminiConfig()._transform_messages( + content = litellm.GoogleAIStudioGeminiConfig().transform_messages( messages=messages, model=model, litellm_params=litellm_params ) else: - content = litellm.VertexGeminiConfig()._transform_messages( + content = litellm.VertexGeminiConfig().transform_messages( messages=messages, model=model, litellm_params=litellm_params ) tools: Final[Tools | None] = optional_params.pop("tools", None) @@ -1237,6 +1245,9 @@ def _transform_request_body( return data +_transform_request_body = transform_request_body + + def sync_transform_request_body( gemini_api_key: str | None, messages: list[AllMessageValues], @@ -1278,7 +1289,7 @@ def sync_transform_request_body( vertex_auth_header=vertex_auth_header, ) - return _transform_request_body( + return transform_request_body( messages=messages, model=model, custom_llm_provider=custom_llm_provider, @@ -1355,7 +1366,7 @@ async def async_transform_request_body( # via _get_gcs_object_content_type to fetch GCS object metadata. Run the # whole sync transformation on a worker thread so it does not block the # async event loop. - return await asyncify(_transform_request_body)( + return await asyncify(transform_request_body)( messages=inlined_messages, model=model, custom_llm_provider=custom_llm_provider, @@ -1364,7 +1375,7 @@ async def async_transform_request_body( optional_params=optional_params, ) - return _transform_request_body( + return transform_request_body( messages=inlined_messages, model=model, custom_llm_provider=custom_llm_provider, @@ -1383,7 +1394,7 @@ def _default_user_message_when_system_message_passed() -> ChatCompletionUserMess return ChatCompletionUserMessage(content=".", role="user") -def _transform_system_message( +def transform_system_message( supports_system_message: bool, messages: list[AllMessageValues] ) -> tuple[SystemInstructions | None, list[AllMessageValues]]: """ @@ -1427,3 +1438,6 @@ def _transform_system_message( return SystemInstructions(parts=system_content_blocks), messages return None, messages + + +_transform_system_message = transform_system_message diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index a291c8e0efb..6bdce393b82 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -34,8 +34,8 @@ from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMExcepti from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.types.llms.anthropic import AnthropicThinkingParam from litellm.types.llms.gemini import BidiGenerateContentServerMessage @@ -89,15 +89,15 @@ from litellm.utils import ( from ....utils import remove_additional_properties, remove_strict_from_schema from ..common_utils import ( VertexAIError, - _build_json_schema, - _build_vertex_schema, + build_json_schema, + build_vertex_schema, supports_response_json_schema, ) from ..vertex_llm_base import VertexBase from .grounding_requests import calculate_grounding_requests from .transformation import ( - _gemini_convert_messages_with_history, async_transform_request_body, + gemini_convert_messages_with_history, sync_transform_request_body, ) @@ -298,6 +298,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False return True + @classmethod + def is_gemini_3_or_newer( + cls, + model: str, + ) -> bool: + return cls._is_gemini_3_or_newer(model) + @staticmethod def _forward_gemini_function_call_id(model: str) -> bool: """ @@ -308,6 +315,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): """ return VertexGeminiConfig._is_gemini_3_or_newer(model) + @classmethod + def forward_gemini_function_call_id( + cls, + model: str, + ) -> bool: + return cls._forward_gemini_function_call_id(model) + def _supports_penalty_parameters(self, model: str) -> bool: # Gemini 3 models do not support penalty parameters if VertexGeminiConfig._is_gemini_3_or_newer(model): @@ -378,6 +392,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): """ return Tools(googleSearch={}) + def map_web_search_options( + self, + value: dict[str, object], # mutable-ok: mirrors override contract + ) -> Tools: + return self._map_web_search_options(value) + @staticmethod def _search_tool_keys() -> set: return { @@ -391,6 +411,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "urlContext", } + @classmethod + def search_tool_keys( + cls, + ) -> set[str]: # mutable-ok: mirrors override contract + return cls._search_tool_keys() + @classmethod def _drop_search_tools_mixed_with_functions(cls, optional_params: dict) -> None: """ @@ -430,6 +456,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): tool for tool in tools if not (isinstance(tool, dict) and any(key in tool for key in search_tool_keys)) ] + @classmethod + def drop_search_tools_mixed_with_functions( + cls, + optional_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> None: + return cls._drop_search_tools_mixed_with_functions(optional_params) + def _map_service_tier_param(self, value: str, optional_params: dict) -> None: """ Map OpenAI service_tier (string) to Gemini serviceTier. @@ -627,7 +660,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): and _openai_function_object["parameters"] is not None and isinstance(_openai_function_object["parameters"], dict) ): # OPENAI accepts JSON Schema, Google accepts OpenAPI schema. - _openai_function_object["parameters"] = _build_vertex_schema(_openai_function_object["parameters"]) + _openai_function_object["parameters"] = build_vertex_schema(_openai_function_object["parameters"]) openai_function_object = _openai_function_object @@ -766,15 +799,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return _tools_list + def map_function( + self, + value: list[dict[str, object]], # mutable-ok: mirrors override contract + optional_params: dict[str, object], # mutable-ok: mirrors override contract + ) -> list[Tools]: # mutable-ok: mirrors override contract + return self._map_function(value, optional_params) + def _map_response_schema(self, value: dict) -> dict: old_schema = deepcopy(value) if isinstance(old_schema, list): for item in old_schema: if isinstance(item, dict): - item = _build_vertex_schema(parameters=item, add_property_ordering=True) + item = build_vertex_schema(parameters=item, add_property_ordering=True) elif isinstance(old_schema, dict): - old_schema = _build_vertex_schema(parameters=old_schema, add_property_ordering=True) + old_schema = build_vertex_schema(parameters=old_schema, add_property_ordering=True) return old_schema def apply_response_schema_transformation(self, value: dict, optional_params: dict, model: str): @@ -813,7 +853,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): # - Standard JSON Schema format (lowercase types) # - Supports additionalProperties # - No propertyOrdering needed - optional_params["response_json_schema"] = _build_json_schema(deepcopy(schema)) + optional_params["response_json_schema"] = build_json_schema(deepcopy(schema)) else: # Use responseSchema (default, backwards compatible) # - OpenAPI-style format (uppercase types) @@ -1065,6 +1105,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return cast(dict, speech_config) + def map_audio_params( + self, + value: dict[str, object], # mutable-ok: mirrors override contract + ) -> dict[str, object]: # mutable-ok: mirrors override contract + return self._map_audio_params(value) + @staticmethod def _apply_include_server_side_tool_invocations( non_default_params: dict, @@ -1299,6 +1345,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return True return False + @classmethod + def is_model_gemini_spec_model( + cls, + model: str | None, + ) -> bool: + return cls._is_model_gemini_spec_model(model) + @staticmethod def _get_model_name_from_gemini_spec_model(model: str) -> str: """ @@ -1917,6 +1970,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return usage + @classmethod + def calculate_usage( + cls, + completion_response: GenerateContentResponseBody | BidiGenerateContentServerMessage, + ) -> Usage: + return cls._calculate_usage(completion_response) + @staticmethod def _check_finish_reason( chat_completion_message: ChatCompletionResponseMessage | None, @@ -1933,6 +1993,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): else: return "stop" + @classmethod + def check_finish_reason( + cls, + chat_completion_message: ChatCompletionResponseMessage | None, + finish_reason: str | None, + ) -> OpenAIChatCompletionFinishReason: + return cls._check_finish_reason(chat_completion_message, finish_reason) + @staticmethod def _check_prompt_level_content_filter( processed_chunk: GenerateContentResponseBody, @@ -1982,10 +2050,26 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return None + @classmethod + def check_prompt_level_content_filter( + cls, + processed_chunk: GenerateContentResponseBody, + response_id: str | None, + model: str | None = None, + ) -> Optional["ModelResponseStream"]: + return cls._check_prompt_level_content_filter(processed_chunk, response_id, model) + @staticmethod def _calculate_web_search_requests(grounding_metadata: list[dict]) -> int | None: return calculate_grounding_requests(grounding_metadata).web_search_requests + @classmethod + def calculate_web_search_requests( + cls, + grounding_metadata: list[dict[str, object]], # mutable-ok: mirrors override contract + ) -> int | None: + return cls._calculate_web_search_requests(grounding_metadata) + @staticmethod def _set_grounding_usage_counters(usage: Usage, grounding_metadata: Sequence[Mapping[str, object]]) -> None: grounding_requests: Final = calculate_grounding_requests(grounding_metadata) @@ -1995,6 +2079,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if grounding_requests.google_maps_grounding_requests is not None: details.google_maps_grounding_requests = grounding_requests.google_maps_grounding_requests + @classmethod + def set_grounding_usage_counters( + cls, + usage: Usage, + grounding_metadata: Sequence[Mapping[str, object]], + ) -> None: + return cls._set_grounding_usage_counters(usage, grounding_metadata) + @staticmethod def _create_streaming_choice( chat_completion_message: ChatCompletionResponseMessage, @@ -2113,6 +2205,19 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): if citation_metadata: model_response._hidden_params["vertex_ai_citation_metadata"] = citation_metadata + @classmethod + def set_stream_metadata_on_response( + cls, + model_response: Union[ModelResponse, "ModelResponseStream"], + grounding_metadata: list[dict[str, object]], # mutable-ok: mirrors override contract + url_context_metadata: list[dict[str, object]], # mutable-ok: mirrors override contract + safety_ratings: list[dict[str, object]], # mutable-ok: mirrors override contract + citation_metadata: list[dict[str, object]], # mutable-ok: mirrors override contract + ) -> None: + return cls._set_stream_metadata_on_response( + model_response, grounding_metadata, url_context_metadata, safety_ratings, citation_metadata + ) + def apply_assembled_streaming_response_metadata( self, response: ModelResponse, @@ -2371,6 +2476,24 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): cumulative_tool_call_index, ) + @classmethod + def process_candidates( + cls, + _candidates: list[Candidates], # mutable-ok: mirrors override contract + model_response: Union[ModelResponse, "ModelResponseStream"], + standard_optional_params: dict[str, object], # mutable-ok: mirrors override contract + cumulative_tool_call_index: int = 0, + ) -> tuple[ # mutable-ok: mirrors override contract + list[dict[str, object]], + list[dict[str, object]], + list[dict[str, object]], + list[dict[str, object]], + int, + ]: + return cls._process_candidates( + _candidates, model_response, standard_optional_params, cumulative_tool_call_index + ) + def transform_response( self, model: str, @@ -2512,19 +2635,39 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return model_response + def transform_google_generate_content_to_openai_model_response( + self, + completion_response: GenerateContentResponseBody | dict[str, object], # mutable-ok: mirrors override contract + model_response: ModelResponse, + model: str, + logging_obj: LoggingClass, + raw_response: httpx.Response, + ) -> ModelResponse: + return self._transform_google_generate_content_to_openai_model_response( + completion_response, model_response, model, logging_obj, raw_response + ) + def _transform_messages( self, messages: list[AllMessageValues], model: str | None = None, litellm_params: dict | None = None, ) -> list[ContentType]: - return _gemini_convert_messages_with_history( + return gemini_convert_messages_with_history( messages=messages, model=model, litellm_params=litellm_params, custom_llm_provider="vertex_ai", ) + def transform_messages( + self, + messages: list[AllMessageValues], # mutable-ok: mirrors override contract + model: str | None = None, + litellm_params: dict[str, object] | None = None, # mutable-ok: mirrors override contract + ) -> list[ContentType]: # mutable-ok: mirrors override contract + return self._transform_messages(messages, model, litellm_params) + def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException: return VertexAIError(message=error_message, status_code=status_code, headers=headers) @@ -3037,7 +3180,7 @@ class VertexLLM(VertexBase): if isinstance(timeout, float) or isinstance(timeout, int): timeout = httpx.Timeout(timeout) _params["timeout"] = timeout - client = _get_httpx_client(params=_params) + client = get_httpx_client(params=_params) else: client = client @@ -3140,7 +3283,7 @@ class ModelResponseIterator: safety_ratings, citation_metadata, self.cumulative_tool_call_index, - ) = VertexGeminiConfig._process_candidates( + ) = VertexGeminiConfig.process_candidates( _candidates, model_response, self.logging_obj.optional_params, @@ -3171,7 +3314,7 @@ class ModelResponseIterator: if self.has_seen_tool_calls: mapped_finish_reason = "tool_calls" else: - mapped_finish_reason = VertexGeminiConfig._check_finish_reason(None, finish_reason_str) + mapped_finish_reason = VertexGeminiConfig.check_finish_reason(None, finish_reason_str) choice = StreamingChoices( finish_reason=mapped_finish_reason, index=candidate.get("index", 0), @@ -3191,7 +3334,7 @@ class ModelResponseIterator: if choice.finish_reason == "stop": choice.finish_reason = "tool_calls" - VertexGeminiConfig._set_stream_metadata_on_response( + VertexGeminiConfig.set_stream_metadata_on_response( model_response, grounding_metadata, url_context_metadata, @@ -3215,11 +3358,11 @@ class ModelResponseIterator: if "usageMetadata" not in processed_chunk: return None - usage: Final = VertexGeminiConfig._calculate_usage( + usage: Final = VertexGeminiConfig.calculate_usage( completion_response=processed_chunk, ) - VertexGeminiConfig._set_grounding_usage_counters(usage, grounding_metadata) + VertexGeminiConfig.set_grounding_usage_counters(usage, grounding_metadata) traffic_type: Final = processed_chunk.get("usageMetadata", {}).get("trafficType") if traffic_type: @@ -3254,7 +3397,7 @@ class ModelResponseIterator: ) # Check if prompt is blocked due to content filtering - blocked_response: Final = VertexGeminiConfig._check_prompt_level_content_filter( + blocked_response: Final = VertexGeminiConfig.check_prompt_level_content_filter( processed_chunk=processed_chunk, response_id=response_id, model=served, diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index c6ac87d646b..d15cf89cf47 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -22,7 +22,7 @@ from litellm.types.utils import EmbeddingResponse from ..gemini.vertex_and_google_ai_studio_gemini import VertexLLM from .batch_embed_content_transformation import ( - _is_file_reference, + is_file_reference, process_embed_content_response, process_response, transform_openai_input_gemini_content, @@ -43,7 +43,7 @@ class GoogleBatchEmbeddings(VertexLLM): flat_elements: Final = [ e for item in input_list for e in (item if isinstance(item, list) else [item]) if isinstance(e, str) ] - has_file_refs: Final = any(_is_file_reference(e) for e in flat_elements) + has_file_refs: Final = any(is_file_reference(e) for e in flat_elements) return flat_elements, has_file_refs def _resolve_file_references( @@ -67,7 +67,7 @@ class GoogleBatchEmbeddings(VertexLLM): resolved_files: Final[dict[str, dict[str, str]]] = {} for element in input_list: - if isinstance(element, str) and _is_file_reference(element): + if isinstance(element, str) and is_file_reference(element): url = f"https://generativelanguage.googleapis.com/v1beta/{element}" headers = {"x-goog-api-key": api_key} response = sync_handler.get(url=url, headers=headers) @@ -104,7 +104,7 @@ class GoogleBatchEmbeddings(VertexLLM): resolved_files: Final[dict[str, dict[str, str]]] = {} for element in input_list: - if isinstance(element, str) and _is_file_reference(element): + if isinstance(element, str) and is_file_reference(element): url = f"https://generativelanguage.googleapis.com/v1beta/{element}" headers = {"x-goog-api-key": api_key} response = await async_handler.get(url=url, headers=headers) diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py index d669acecfd9..9ee2b71d30c 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_transformation.py @@ -40,11 +40,14 @@ SUPPORTED_EMBEDDING_MIME_TYPES: Final = { } -def _is_file_reference(s: str) -> bool: +def is_file_reference(s: str) -> bool: """Check if string is a Gemini file reference (files/...).""" return isinstance(s, str) and s.startswith("files/") +_is_file_reference = is_file_reference + + def _is_gcs_url(s: str) -> bool: """Check if string is a GCS URL (gs://...).""" return isinstance(s, str) and s.startswith("gs://") @@ -152,7 +155,7 @@ def _is_multimodal_element(element: str) -> bool: """Check if a single string element is multimodal.""" if element.startswith("data:") and ";base64," in element: return True - if _is_file_reference(element): + if is_file_reference(element): return True if _is_gcs_url(element): return True @@ -180,7 +183,7 @@ def _build_part_for_input( "file_uri": element, } return PartType(file_data=file_data) - elif _is_file_reference(element): + elif is_file_reference(element): if element not in resolved_files: raise ValueError(f"File reference {element} not resolved") file_info: Final = resolved_files[element] @@ -331,7 +334,7 @@ def _is_image_element( return _infer_mime_type_from_gcs_url(element) in _IMAGE_MIME_TYPES except ValueError: return False - if _is_file_reference(element): + if is_file_reference(element): file_info: Final = resolved_files.get(element) return file_info is not None and file_info.get("mime_type") in _IMAGE_MIME_TYPES return False diff --git a/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py b/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py index 7975e708428..d053f564dee 100644 --- a/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py +++ b/litellm/llms/vertex_ai/text_to_speech/text_to_speech_handler.py @@ -5,8 +5,8 @@ from typing_extensions import TypedDict import litellm from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.llms.openai.openai import HttpxBinaryResponseContent from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM @@ -140,7 +140,7 @@ class VertexTextToSpeechAPI(VertexLLM): ####### Send the request ################### if _is_async is True: return self.async_audio_speech(logging_obj=logging_obj, url=url, headers=headers, request=request) - sync_handler: Final = _get_httpx_client() + sync_handler: Final = get_httpx_client() response = sync_handler.post( url=url, diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py index be2dacd23c2..5d748ad6e74 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py @@ -3,7 +3,7 @@ from typing import Any, Final from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.anthropic.pass_through.messages.transformation import ( AnthropicMessagesConfig, - _messages_carry_output_config, + messages_carry_output_config, ) from litellm.types.llms.anthropic import ( ANTHROPIC_BETA_HEADER_VALUES, @@ -112,7 +112,7 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert if optional_params.get("safeguards") is not None: beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.DANGEROUS_TOOL_USE_2026_09_03.value) - if _messages_carry_output_config(messages): + if messages_carry_output_config(messages): beta_values.add(ANTHROPIC_BETA_HEADER_VALUES.PER_TURN_CONTROL_2026_07_01.value) if beta_values: diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/output_params_utils.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/output_params_utils.py index 4c7f93eb437..64ba1ad21cb 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/output_params_utils.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/output_params_utils.py @@ -28,7 +28,7 @@ def _model_accepts_output_config_effort(model: str) -> bool: """ from litellm.llms.anthropic.chat.transformation import AnthropicConfig - return AnthropicConfig._model_supports_effort_param(model, "vertex_ai") + return AnthropicConfig.model_supports_effort_param(model, "vertex_ai") def sanitize_vertex_anthropic_output_params(data: dict, model: str) -> None: diff --git a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py index 15378839b33..9626c879dd0 100644 --- a/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py +++ b/litellm/llms/vertex_ai/vertex_embeddings/embedding_handler.py @@ -7,8 +7,8 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.llms.vertex_ai.vertex_ai_non_gemini import VertexAIError from litellm.llms.vertex_ai.vertex_llm_base import VertexBase @@ -100,7 +100,7 @@ class VertexEmbedding(VertexBase): if timeout: _client_params["timeout"] = timeout if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params=_client_params) + client = get_httpx_client(params=_client_params) else: client = client ## LOGGING diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index 08d95430041..9fa86ec30f5 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -18,7 +18,7 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, - _get_httpx_client, + get_httpx_client, ) from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.llms.vertex_ai.common_utils import VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS @@ -192,7 +192,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): json=request_data, timeout=timeout, ) - return _get_httpx_client().post( + return get_httpx_client().post( url=api_base, headers=headers, json=request_data, diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 8b7f8c63625..1d5cbe09ba0 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -20,15 +20,15 @@ from litellm.secret_managers.main import get_secret_str from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES, VertexPartnerProvider from .common_utils import ( - _get_gemini_url, - _get_vertex_url, all_gemini_url_modes, + get_gemini_url, get_vertex_base_model_name, get_vertex_base_url, + get_vertex_url, ) -def _graft_default_vertex_path(api_base: str, default_url: str) -> str: +def graft_default_vertex_path(api_base: str, default_url: str) -> str: parsed_api_base: Final = urlparse(api_base) default_segments: Final = urlparse(default_url).path.lstrip("/").split("/") graft_segments: Final = default_segments[1:] if default_segments[0] in ("v1", "v1beta1") else default_segments @@ -36,6 +36,8 @@ def _graft_default_vertex_path(api_base: str, default_url: str) -> str: return parsed_api_base._replace(path=grafted_path).geturl() +_graft_default_vertex_path = graft_default_vertex_path + GOOGLE_IMPORT_ERROR_MESSAGE: Final = ( "Google Cloud SDK not found. Install it with: pip install 'litellm[google]' or pip install google-cloud-aiplatform" ) @@ -618,6 +620,14 @@ class VertexBase: project_id=project_id, ) + def ensure_access_token( + self, + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], + ) -> tuple[str, str]: + return self._ensure_access_token(credentials, project_id, custom_llm_provider) + def _check_custom_proxy( self, api_base: str | None, @@ -688,7 +698,7 @@ class VertexBase: elif urlparse(api_base).path in ("", "/"): url = api_base.rstrip("/") + urlparse(url).path elif urlparse(api_base).path.rstrip("/") in ("/v1", "/v1beta1") and "/projects/" in urlparse(url).path: - url = _graft_default_vertex_path(api_base=api_base, default_url=url) + url = graft_default_vertex_path(api_base=api_base, default_url=url) else: url = f"{api_base}:{endpoint}" if stream is True: @@ -726,7 +736,7 @@ class VertexBase: raise ValueError( "Missing Gemini API key. Set the GEMINI_API_KEY or GOOGLE_API_KEY environment variable." ) - url, endpoint = _get_gemini_url( + url, endpoint = get_gemini_url( mode=mode, model=model, stream=stream, @@ -740,7 +750,7 @@ class VertexBase: ### SET RUNTIME ENDPOINT ### version = "v1beta1" if should_use_v1beta1_features is True else "v1" - url, endpoint = _get_vertex_url( + url, endpoint = get_vertex_url( mode=mode, model=model, stream=stream, @@ -764,6 +774,36 @@ class VertexBase: use_psc_endpoint_format=use_psc_endpoint_format, ) + def get_token_and_url( + self, + model: str, + auth_header: str | None, + gemini_api_key: str | None, + vertex_project: str | None, + vertex_location: str | None, + vertex_credentials: VERTEX_CREDENTIALS_TYPES | None, + stream: bool | None, + custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], + api_base: str | None, + should_use_v1beta1_features: bool | None = False, + mode: all_gemini_url_modes = "chat", + use_psc_endpoint_format: bool = False, + ) -> tuple[str | None, str]: + return self._get_token_and_url( + model, + auth_header, + gemini_api_key, + vertex_project, + vertex_location, + vertex_credentials, + stream, + custom_llm_provider, + api_base, + should_use_v1beta1_features, + mode, + use_psc_endpoint_format, + ) + def _handle_reauthentication( self, credentials: VERTEX_CREDENTIALS_TYPES | None, @@ -1135,6 +1175,14 @@ class VertexBase: project_id=project_id, ) + async def ensure_access_token_async( + self, + credentials: VERTEX_CREDENTIALS_TYPES | None, + project_id: str | None, + custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], + ) -> tuple[str, str]: + return await self._ensure_access_token_async(credentials, project_id, custom_llm_provider) + def set_headers(self, auth_header: str | None, extra_headers: dict | None) -> dict: headers: Final = { "Content-Type": "application/json", diff --git a/litellm/llms/vertex_ai/videos/transformation.py b/litellm/llms/vertex_ai/videos/transformation.py index 6c29059f0a9..295fe8a7350 100644 --- a/litellm/llms/vertex_ai/videos/transformation.py +++ b/litellm/llms/vertex_ai/videos/transformation.py @@ -20,7 +20,7 @@ from litellm.constants import DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS from litellm.images.utils import ImageEditRequestUtils from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.vertex_ai.common_utils import ( - _convert_vertex_datetime_to_openai_datetime, + convert_vertex_datetime_to_openai_datetime, get_vertex_base_url, ) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase @@ -537,7 +537,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase): create_time_str: Final = response_data.get("metadata", {}).get("createTime") if create_time_str: try: - created_at = _convert_vertex_datetime_to_openai_datetime(create_time_str) + created_at = convert_vertex_datetime_to_openai_datetime(create_time_str) except Exception: created_at = int(time.time()) else: diff --git a/litellm/llms/watsonx/chat/handler.py b/litellm/llms/watsonx/chat/handler.py index da466e8c62a..0f565cbb36f 100644 --- a/litellm/llms/watsonx/chat/handler.py +++ b/litellm/llms/watsonx/chat/handler.py @@ -7,7 +7,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.utils import CustomStreamingDecoder, ModelResponse from ...openai_like.chat.handler import OpenAILikeChatHandler -from ..common_utils import _get_api_params +from ..common_utils import get_api_params from .transformation import IBMWatsonXChatConfig watsonx_chat_transformation: Final = IBMWatsonXChatConfig() @@ -41,7 +41,7 @@ class WatsonXChatHandler(OpenAILikeChatHandler): streaming_decoder: CustomStreamingDecoder | None = None, fake_stream: bool = False, ): - api_params: Final = _get_api_params(params=optional_params, model=model) + api_params: Final = get_api_params(params=optional_params, model=model) ## UPDATE HEADERS headers = watsonx_chat_transformation.validate_environment( @@ -54,7 +54,7 @@ class WatsonXChatHandler(OpenAILikeChatHandler): ) ## UPDATE PAYLOAD (optional params and special cases for models deployed in spaces) - watsonx_auth_payload: Final = watsonx_chat_transformation._prepare_payload( + watsonx_auth_payload: Final = watsonx_chat_transformation.prepare_payload( model=model, api_params=api_params, ) diff --git a/litellm/llms/watsonx/common_utils.py b/litellm/llms/watsonx/common_utils.py index 7b567e4fab1..fb1fa8b9ffb 100644 --- a/litellm/llms/watsonx/common_utils.py +++ b/litellm/llms/watsonx/common_utils.py @@ -70,14 +70,17 @@ def generate_iam_token(api_key=None, **params) -> str: return cast(str, result) -def _generate_watsonx_token(api_key: str | None, token: str | None) -> str: +def generate_watsonx_token(api_key: str | None, token: str | None) -> str: if token is not None: return token token = generate_iam_token(api_key) return token -def _get_api_params(params: dict, model: str | None = None) -> WatsonXAPIParams: +_generate_watsonx_token = generate_watsonx_token + + +def get_api_params(params: dict, model: str | None = None) -> WatsonXAPIParams: """ Find watsonx.ai credentials in the params or environment variables and return the headers for authentication. """ @@ -120,6 +123,9 @@ def _get_api_params(params: dict, model: str | None = None) -> WatsonXAPIParams: ) +_get_api_params = get_api_params + + async def _aconvert_watsonx_messages_core( model: str, messages: list[AllMessageValues], @@ -252,7 +258,7 @@ class IBMWatsonXMixin: elif zen_api_key: headers["Authorization"] = f"ZenApiKey {zen_api_key}" else: - token = _generate_watsonx_token(api_key=api_key, token=token) + token = generate_watsonx_token(api_key=api_key, token=token) # build auth headers headers["Authorization"] = f"Bearer {token}" return {**default_headers, **headers} @@ -343,3 +349,10 @@ class IBMWatsonXMixin: else: payload["space_id"] = api_params["space_id"] return payload + + def prepare_payload( + self, + model: str, + api_params: WatsonXAPIParams, + ) -> dict[str, object]: # mutable-ok: mirrors override contract + return self._prepare_payload(model, api_params) diff --git a/litellm/llms/watsonx/completion/transformation.py b/litellm/llms/watsonx/completion/transformation.py index 2fec8485cf9..e01a0ff85e5 100644 --- a/litellm/llms/watsonx/completion/transformation.py +++ b/litellm/llms/watsonx/completion/transformation.py @@ -15,9 +15,9 @@ from ...base_llm.chat.transformation import BaseConfig from ..common_utils import ( IBMWatsonXMixin, WatsonXAIError, - _get_api_params, aconvert_watsonx_messages_to_prompt, convert_watsonx_messages_to_prompt, + get_api_params, ) if TYPE_CHECKING: @@ -226,7 +226,7 @@ class IBMWatsonXAIConfig(IBMWatsonXMixin, BaseConfig): """Shared logic to build request payload""" extra_body_params: Final = optional_params.pop("extra_body", {}) optional_params.update(extra_body_params) - watsonx_api_params: Final = _get_api_params(params=optional_params, model=model) + watsonx_api_params: Final = get_api_params(params=optional_params, model=model) watsonx_auth_payload: Final = self._prepare_payload(model=model, api_params=watsonx_api_params) return { diff --git a/litellm/llms/watsonx/embed/transformation.py b/litellm/llms/watsonx/embed/transformation.py index 609eeba90c0..2e1e8fac643 100644 --- a/litellm/llms/watsonx/embed/transformation.py +++ b/litellm/llms/watsonx/embed/transformation.py @@ -16,7 +16,7 @@ from litellm.types.llms.openai import AllEmbeddingInputValues from litellm.types.llms.watsonx import WatsonXAIEndpoint from litellm.types.utils import EmbeddingResponse, Usage -from ..common_utils import IBMWatsonXMixin, _get_api_params +from ..common_utils import IBMWatsonXMixin, get_api_params _JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) _TOKEN_COUNT: Final = TypeAdapter(int) @@ -42,7 +42,7 @@ class IBMWatsonXEmbeddingConfig(IBMWatsonXMixin, BaseEmbeddingConfig): optional_params: dict, headers: dict, ) -> dict: - watsonx_api_params: Final = _get_api_params(params=optional_params, model=model) + watsonx_api_params: Final = get_api_params(params=optional_params, model=model) watsonx_auth_payload: Final = self._prepare_payload( model=model, api_params=watsonx_api_params, diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index c2cfd6597f3..5f337603763 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -23,7 +23,7 @@ from litellm.types.rerank import ( RerankTokens, ) -from ..common_utils import IBMWatsonXMixin, _generate_watsonx_token, _get_api_params +from ..common_utils import IBMWatsonXMixin, generate_watsonx_token, get_api_params _JSON_OBJECT: Final = TypeAdapter(Mapping[str, object], config=ConfigDict(hide_input_in_errors=True)) _JSON_OBJECTS: Final = TypeAdapter(Iterable[Mapping[str, object]], config=ConfigDict(hide_input_in_errors=True)) @@ -91,7 +91,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): elif zen_api_key: headers["Authorization"] = f"ZenApiKey {zen_api_key}" else: - token = _generate_watsonx_token(api_key=api_key, token=token) + token = generate_watsonx_token(api_key=api_key, token=token) # build auth headers headers["Authorization"] = f"Bearer {token}" return {**default_headers, **headers} @@ -146,7 +146,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): """ Transform request to IBM watsonx.ai rerank format """ - watsonx_api_params: Final = _get_api_params(params=optional_rerank_params, model=model) + watsonx_api_params: Final = get_api_params(params=optional_rerank_params, model=model) watsonx_auth_payload: Final = self._prepare_payload( model=model, api_params=watsonx_api_params, diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index 86073eb672d..984b62eeb59 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -53,6 +53,13 @@ class XAIChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = XAIModelInfo.get_api_key(api_key) return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def validate_environment( self, headers: dict, @@ -167,6 +174,12 @@ class XAIChatConfig(OpenAIGPTConfig): return False return True + def supports_stop_reason( + self, + model: str, + ) -> bool: + return self._supports_stop_reason(model) + def _supports_frequency_penalty(self, model: str) -> bool: """ From manual testing grok-4 does not support `frequency_penalty` @@ -406,6 +419,13 @@ class XAIChatConfig(OpenAIGPTConfig): if int(usage.total_tokens or 0) < expected_total: usage.total_tokens = expected_total + @classmethod + def normalize_openai_compatible_usage_totals( + cls, + usage: Usage | dict[str, object] | None, # mutable-ok: mirrors override contract + ) -> None: + return cls._normalize_openai_compatible_usage_totals(usage) + class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler): def chunk_parser(self, chunk: dict) -> ModelResponseStream: @@ -428,7 +448,7 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler): if "usage" in chunk and chunk["usage"] is not None: XAIChatConfig.fold_reasoning_tokens_into_completion(chunk["usage"]) - XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"]) + XAIChatConfig.normalize_openai_compatible_usage_totals(chunk["usage"]) parsed_chunk: Final = super().chunk_parser(chunk) restated_usage: Final = _usage_restated_from_xai_ticks(getattr(parsed_chunk, "usage", None)) diff --git a/litellm/llms/xai/oauth.py b/litellm/llms/xai/oauth.py index e8196ec6cb9..56dcc5a0c90 100644 --- a/litellm/llms/xai/oauth.py +++ b/litellm/llms/xai/oauth.py @@ -18,7 +18,7 @@ from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm._logging import verbose_logger from litellm.constants import XAI_API_BASE -from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client +from litellm.llms.custom_httpx.http_handler import HTTPHandler, get_httpx_client from litellm.secret_managers.main import get_secret_str XAI_OAUTH_ISSUER: Final = "https://auth.x.ai" @@ -204,7 +204,7 @@ class XAIOAuthAuthenticator: return auth_data def _client(self) -> httpx.Client | HTTPHandler: - return self.http_client or _get_httpx_client() + return self.http_client or get_httpx_client() def _ensure_token_dir(self) -> None: os.makedirs(self.token_dir, mode=0o700, exist_ok=True) diff --git a/litellm/llms/zai/chat/transformation.py b/litellm/llms/zai/chat/transformation.py index b53eef19d1e..e0e42d772f8 100644 --- a/litellm/llms/zai/chat/transformation.py +++ b/litellm/llms/zai/chat/transformation.py @@ -20,6 +20,13 @@ class ZAIChatConfig(OpenAIGPTConfig): dynamic_api_key: Final = api_key or get_secret_str("ZAI_API_KEY") return api_base, dynamic_api_key + def get_openai_compatible_provider_info( + self, + api_base: str | None, + api_key: str | None, + ) -> tuple[str | None, str | None]: + return self._get_openai_compatible_provider_info(api_base, api_key) + def remove_cache_control_flag_from_messages_and_tools( self, model: str, diff --git a/litellm/main.py b/litellm/main.py index 3736051459d..519b8aad0ef 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -211,7 +211,7 @@ from .litellm_core_utils.prompt_templates.factory import ( from .litellm_core_utils.streaming_chunk_builder_utils import ChunkProcessor from .llms.anthropic.chat import AnthropicChatCompletion from .llms.azure.audio_transcriptions import AzureAudioTranscription -from .llms.azure.azure import AzureChatCompletion, _check_dynamic_azure_params +from .llms.azure.azure import AzureChatCompletion, check_dynamic_azure_params from .llms.azure.chat.o_series_handler import AzureOpenAIO1ChatCompletion from .llms.azure.completion.handler import AzureTextCompletion from .llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion @@ -1361,7 +1361,7 @@ def _complete_azure(ctx: CompletionDispatchContext) -> _CompletionDispatchResult dynamic_params = False if client is not None and (isinstance(client, openai.AzureOpenAI) or isinstance(client, openai.AsyncAzureOpenAI)): - dynamic_params = _check_dynamic_azure_params( + dynamic_params = check_dynamic_azure_params( azure_client_params={"api_version": api_version}, azure_client=client, ) @@ -5068,7 +5068,7 @@ def _complete_langgraph(ctx: CompletionDispatchContext) -> _CompletionDispatchRe ( api_base, api_key, - ) = LangGraphConfig()._get_openai_compatible_provider_info( + ) = LangGraphConfig().get_openai_compatible_provider_info( api_base=api_base or litellm.api_base, api_key=api_key or litellm.api_key, ) @@ -5117,7 +5117,7 @@ def _complete_langflow(ctx: CompletionDispatchContext) -> _CompletionDispatchRes ( api_base, api_key, - ) = LangFlowConfig()._get_openai_compatible_provider_info( + ) = LangFlowConfig().get_openai_compatible_provider_info( api_base=api_base or litellm.api_base, api_key=api_key or litellm.api_key, ) @@ -7846,7 +7846,7 @@ async def amoderation( if openai_client is None or not isinstance(openai_client, AsyncOpenAI): # call helper to get OpenAI client # _get_openai_client maintains in-memory caching logic for OpenAI clients - _openai_client: AsyncOpenAI = openai_chat_completions._get_openai_client( + _openai_client: AsyncOpenAI = openai_chat_completions.get_openai_client( is_async=True, api_key=api_key, api_base=optional_params.api_base or _dynamic_api_base, diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 17b97dee37f..ff9eb8fe79a 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -383,7 +383,7 @@ async def allm_passthrough_route( # If no provider config available, raise the original exception raise e - raise base_llm_http_handler._handle_error( + raise base_llm_http_handler.handle_error( e=e, provider_config=provider_config, ) @@ -441,8 +441,8 @@ def llm_passthrough_route( if client is None: from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.passthrough.timeout_utils import resolve_llm_passthrough_timeout from litellm.types.llms.custom_http import httpxSpecialProvider @@ -457,7 +457,7 @@ def llm_passthrough_route( params={"timeout": resolved_timeout}, ) else: - client = _get_httpx_client(params={"timeout": resolved_timeout}) + client = get_httpx_client(params={"timeout": resolved_timeout}) # Add model_id to litellm_params if present in kwargs (for Bedrock Application Inference Profiles) if "model_id" in kwargs: @@ -595,7 +595,7 @@ def llm_passthrough_route( return response except Exception as e: assert provider_config is not None - raise base_llm_http_handler._handle_error( + raise base_llm_http_handler.handle_error( e=e, provider_config=provider_config, ) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index fceae2d929c..ff12f75479a 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -1783,7 +1783,7 @@ async def _get_health_readiness_details( "cache": cache_type, "litellm_version": version, "success_callbacks": success_callback_names, - "use_aiohttp_transport": AsyncHTTPHandler._should_use_aiohttp_transport(), + "use_aiohttp_transport": AsyncHTTPHandler.should_use_aiohttp_transport(), "log_level": log_level_name, "is_detailed_debug": is_detailed_debug, "show_no_redis_warning": show_no_redis_warning, @@ -1796,7 +1796,7 @@ async def _get_health_readiness_details( "cache": cache_type, "litellm_version": version, "success_callbacks": success_callback_names, - "use_aiohttp_transport": AsyncHTTPHandler._should_use_aiohttp_transport(), + "use_aiohttp_transport": AsyncHTTPHandler.should_use_aiohttp_transport(), "log_level": log_level_name, "is_detailed_debug": is_detailed_debug, "show_no_redis_warning": show_no_redis_warning, diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 88c8a67568b..46f2394c845 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -2726,13 +2726,13 @@ async def _prepare_vertex_auth_headers( else: raise ValueError("No vertex credentials found") - _auth_header, vertex_project = await vertex_llm_base._ensure_access_token_async( + _auth_header, vertex_project = await vertex_llm_base.ensure_access_token_async( credentials=vertex_credentials_str, project_id=vertex_project, custom_llm_provider="vertex_ai_beta", ) - auth_header, _ = vertex_llm_base._get_token_and_url( + auth_header, _ = vertex_llm_base.get_token_and_url( model="", auth_header=_auth_header, gemini_api_key=None, @@ -3771,7 +3771,7 @@ async def vertex_ai_live_websocket_passthrough( ( access_token, resolved_project, - ) = await vertex_llm_base._ensure_access_token_async( + ) = await vertex_llm_base.ensure_access_token_async( credentials=credentials_value, project_id=configured_project, custom_llm_provider="vertex_ai_beta", diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py index 8bba698ec71..a11b0be3983 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/cohere_passthrough_logging_handler.py @@ -95,7 +95,7 @@ class CoherePassthroughLoggingHandler(BasePassthroughLoggingHandler): input_texts = request_body.get("input", []) # Transform the response - litellm_model_response = cohere_embed_config._transform_response( + litellm_model_response = cohere_embed_config.transform_response( response=httpx_response, api_key="", logging_obj=logging_obj, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index 89c21d092ff..0fcc1ccda38 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -578,7 +578,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): ) # Convert string chunk to dict - stripped_json_chunk = BaseModelResponseIterator._string_to_dict_parser(str_line=chunk_str) + stripped_json_chunk = BaseModelResponseIterator.string_to_dict_parser(str_line=chunk_str) if stripped_json_chunk: # Parse the chunk using OpenAI's chunk parser diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index e749c2b992e..47e1f21d1a9 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -644,7 +644,7 @@ class VertexPassthroughLoggingHandler: ) chunk_parsing_logic = vertex_iterator.chunk_parser for chunk in all_chunks: - dict_chunk = BaseModelResponseIterator._string_to_dict_parser(chunk) + dict_chunk = BaseModelResponseIterator.string_to_dict_parser(chunk) if dict_chunk is None: continue parsed_chunks.append(chunk_parsing_logic(dict_chunk)) @@ -824,7 +824,7 @@ class VertexPassthroughLoggingHandler: ) # Extract batch ID and model from the response - batch_id = VertexAIBatchTransformation._get_batch_id_from_vertex_ai_batch_response(_json_response) + batch_id = VertexAIBatchTransformation.get_batch_id_from_vertex_ai_batch_response(_json_response) model_name: Final = _json_response.get("model", "unknown") # Create unified object ID for tracking diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py index 9a259b4778a..8199a96ceab 100644 --- a/litellm/proxy/pass_through_endpoints/streaming_handler.py +++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py @@ -295,8 +295,8 @@ class PassThroughStreamingHandler: - OpenAI """ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( - _is_message_stop_chunk, # pyright: ignore[reportPrivateUsage] # both native stream paths share terminal-event detection - _is_provider_error_chunk, # pyright: ignore[reportPrivateUsage] # provider errors must not become cache evidence + is_message_stop_chunk, # pyright: ignore[reportPrivateUsage] # both native stream paths share terminal-event detection + is_provider_error_chunk, # pyright: ignore[reportPrivateUsage] # provider errors must not become cache evidence ) # Transport reads can split event names and JSON payloads. Recognize terminal @@ -309,8 +309,8 @@ class PassThroughStreamingHandler: ] = ( endpoint_type == EndpointType.ANTHROPIC and not incomplete_tail.strip() - and _is_message_stop_chunk(complete_frames) - and not _is_provider_error_chunk(complete_frames) + and is_message_stop_chunk(complete_frames) + and not is_provider_error_chunk(complete_frames) ) try: # TinyFish billing is owned by the detached poller; the $0 fallback below is only for streams with no run_id diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e965a872ad9..db6294a364d 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1233,7 +1233,7 @@ async def _initialize_shared_aiohttp_session(): from aiohttp import ClientSession, DummyCookieJar, TCPConnector from litellm.llms.custom_httpx.http_handler import ( - _build_aiohttp_keepalive_socket_factory, + build_aiohttp_keepalive_socket_factory, ) connector_kwargs: Final[_AiohttpConnectorKwargs] = { @@ -1246,7 +1246,7 @@ async def _initialize_shared_aiohttp_session(): connector_kwargs["limit"] = AIOHTTP_CONNECTOR_LIMIT if AIOHTTP_CONNECTOR_LIMIT_PER_HOST > 0: connector_kwargs["limit_per_host"] = AIOHTTP_CONNECTOR_LIMIT_PER_HOST - socket_factory: Final = _build_aiohttp_keepalive_socket_factory() + socket_factory: Final = build_aiohttp_keepalive_socket_factory() if socket_factory is not None: connector_kwargs["socket_factory"] = socket_factory diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index eeee04f82bb..83cacf6c4d9 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -315,7 +315,7 @@ async def vertex_access_token_resolver( project_id: str | None, custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], ) -> tuple[str, str]: - return await vertex_llm_base._ensure_access_token_async( + return await vertex_llm_base.ensure_access_token_async( credentials=credentials, project_id=project_id, custom_llm_provider=custom_llm_provider, @@ -679,7 +679,7 @@ async def realtime_health_check( realtime_protocol=realtime_protocol, model_params=resolved_params, ) - url = azure_realtime._construct_url( + url = azure_realtime.construct_url( api_base=resolved_api_base or "", model=model, api_version=resolved_api_version or "2024-10-01-preview", @@ -687,12 +687,12 @@ async def realtime_health_check( query_params=azure_query_params, ) elif custom_llm_provider == "openai": - url = openai_realtime._construct_url( + url = openai_realtime.construct_url( api_base=resolved_api_base or "https://api.openai.com/", query_params={"model": model}, ) elif custom_llm_provider == "xai": - url = xai_realtime._construct_url( + url = xai_realtime.construct_url( api_base=resolved_api_base or "https://api.x.ai/v1", query_params={"model": model} ) elif custom_llm_provider == "vertex_ai": diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index 1bd4a4fa566..a702c206cfd 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -27,8 +27,8 @@ from litellm._logging import verbose_logger from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, ) from litellm.secret_managers.main import get_secret_str from litellm.types.llms.custom_http import httpxSpecialProvider @@ -197,7 +197,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): optional_params=optional_params, ) - sync_client: Final = _get_httpx_client( + sync_client: Final = get_httpx_client( params={"timeout": timeout}, ) diff --git a/litellm/secret_managers/cyberark_secret_manager.py b/litellm/secret_managers/cyberark_secret_manager.py index 2c827e36979..a1ec7c0a386 100644 --- a/litellm/secret_managers/cyberark_secret_manager.py +++ b/litellm/secret_managers/cyberark_secret_manager.py @@ -12,8 +12,8 @@ from litellm._logging import verbose_logger from litellm.caching import InMemoryCache from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) from litellm.types.secret_managers.main import KeyManagementSystem @@ -95,7 +95,7 @@ class CyberArkSecretManager(BaseSecretManager): resp = http_client.post(auth_url, content=self.conjur_api_key) else: # API key authentication - http_handler: Final = _get_httpx_client(params={"ssl_verify": self.ssl_verify}) + http_handler: Final = get_httpx_client(params={"ssl_verify": self.ssl_verify}) resp = http_handler.client.post(auth_url, content=self.conjur_api_key) resp.raise_for_status() @@ -248,7 +248,7 @@ class CyberArkSecretManager(BaseSecretManager): if self.cache.get_cache(secret_name) is not None: return self.cache.get_cache(secret_name) - sync_client: Final = _get_httpx_client(params={"ssl_verify": self.ssl_verify}) + sync_client: Final = get_httpx_client(params={"ssl_verify": self.ssl_verify}) try: url: Final = self.get_url(secret_name) diff --git a/litellm/secret_managers/google_secret_manager.py b/litellm/secret_managers/google_secret_manager.py index a29e5ab96e1..2a5152179bb 100644 --- a/litellm/secret_managers/google_secret_manager.py +++ b/litellm/secret_managers/google_secret_manager.py @@ -7,7 +7,7 @@ from litellm._logging import verbose_logger from litellm.caching.caching import InMemoryCache from litellm.constants import SECRET_MANAGER_REFRESH_INTERVAL from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase -from litellm.llms.custom_httpx.http_handler import _get_httpx_client +from litellm.llms.custom_httpx.http_handler import get_httpx_client from litellm.proxy._types import CommonProxyErrors from litellm.types.secret_managers.main import KeyManagementSystem @@ -35,7 +35,7 @@ class GoogleSecretManager(GCSBucketBase): raise ValueError( "Google Secret Manager requires a project ID, please set 'GOOGLE_SECRET_MANAGER_PROJECT_ID' in your .env" ) - self.sync_httpx_client = _get_httpx_client() + self.sync_httpx_client = get_httpx_client() litellm.secret_manager_client = self litellm._key_management_system = KeyManagementSystem.GOOGLE_SECRET_MANAGER _refresh_interval = os.environ.get("GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL", refresh_interval) diff --git a/litellm/secret_managers/hashicorp_secret_manager.py b/litellm/secret_managers/hashicorp_secret_manager.py index 5a05ba8f965..343492b216a 100644 --- a/litellm/secret_managers/hashicorp_secret_manager.py +++ b/litellm/secret_managers/hashicorp_secret_manager.py @@ -11,8 +11,8 @@ from litellm._logging import verbose_logger from litellm.caching import InMemoryCache from litellm.constants import SECRET_MANAGER_REFRESH_INTERVAL from litellm.llms.custom_httpx.http_handler import ( - _get_httpx_client, get_async_httpx_client, + get_httpx_client, httpxSpecialProvider, ) from litellm.types.secret_managers.main import KeyManagementSystem @@ -191,7 +191,7 @@ class HashicorpSecretManager(BaseSecretManager): headers: Final = self._get_login_headers() try: - client: Final = _get_httpx_client() + client: Final = get_httpx_client() resp: Final = client.post( url=login_url, headers=headers, @@ -436,7 +436,7 @@ class HashicorpSecretManager(BaseSecretManager): secret_name is just the path inside the KV mount (e.g., 'myapp/config'). Returns the entire data dict from data.data, or None on failure. """ - sync_client: Final = _get_httpx_client() + sync_client: Final = get_httpx_client() try: target: Final = self._build_secret_target(secret_name, optional_params) cached_body: Final = self.cache.get_cache(target["url"]) diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index ca0cb7bd132..0c53137cfab 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -13,7 +13,7 @@ IGNORE_FUNCTIONS = [ "convert_anyof_null_to_nullable", # has a set max depth "add_object_type", "strip_field", - "_transform_prompt", + "transform_prompt", "mask_dict", "_serialize", # we now set a max depth for this "_sanitize_request_body_for_spend_logs_payload", # testing added for circular reference diff --git a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py index 1be3ca0745d..4fc79836e2e 100644 --- a/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py +++ b/tests/image_gen_tests/test_bedrock_image_gen_unit_tests.py @@ -50,7 +50,7 @@ from litellm.llms.bedrock.common_utils import BedrockError ], ) def test_is_stability_3_model(model, expected): - result = AmazonStability3Config._is_stability_3_model(model) + result = AmazonStability3Config.is_stability_3_model(model) assert result == expected @@ -69,7 +69,7 @@ def test_is_stability_3_model(model, expected): ], ) def test_is_nova_canvas_model(model, expected): - result = AmazonNovaCanvasConfig._is_nova_model(model) + result = AmazonNovaCanvasConfig.is_nova_model(model) assert result == expected @@ -292,9 +292,7 @@ def test_transform_request_body_with_invalid_task_type(): optional_params = {"taskType": "INVALID_TASK"} with pytest.raises(NotImplementedError) as exc_info: - AmazonNovaCanvasConfig.transform_request_body( - text=text, optional_params=optional_params - ) + AmazonNovaCanvasConfig.transform_request_body(text=text, optional_params=optional_params) assert "Task type INVALID_TASK is not supported" in str(exc_info.value) diff --git a/tests/litellm_utils_tests/test_cyberark.py b/tests/litellm_utils_tests/test_cyberark.py index fefc93af724..84b1fc5ebd2 100644 --- a/tests/litellm_utils_tests/test_cyberark.py +++ b/tests/litellm_utils_tests/test_cyberark.py @@ -55,7 +55,7 @@ async def test_cyberark_write_secret_rejects_yaml_injection(): with ( patch( - "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", return_value=mock_sync_client, ), patch( @@ -106,7 +106,7 @@ async def test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters(secre mock_async_client.client.post.side_effect = _capture_post with patch( - "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", return_value=mock_sync_client, ): cyberark_manager = CyberArkSecretManager() @@ -131,26 +131,20 @@ async def test_cyberark_write_and_read_secret(): secret_value = f"test-value-{uuid.uuid4()}" # Mock sync httpx client (for auth, ensure variable exists, sync read) - # The _get_httpx_client returns an HTTPHandler with a .client property + # The get_httpx_client returns an HTTPHandler with a .client property mock_sync_client = MagicMock() # Auth response - note: the actual client is accessed via .client property - mock_sync_client.client.post.return_value = create_mock_response( - status_code=200, text="mock-token" - ) + mock_sync_client.client.post.return_value = create_mock_response(status_code=200, text="mock-token") # Sync read response - mock_sync_client.client.get.return_value = create_mock_response( - status_code=200, text=secret_value - ) + mock_sync_client.client.get.return_value = create_mock_response(status_code=200, text=secret_value) # Mock async httpx client (for async write) mock_async_client = AsyncMock() - mock_async_client.post.return_value = create_mock_response( - status_code=201, text="" - ) + mock_async_client.post.return_value = create_mock_response(status_code=201, text="") with ( patch( - "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", return_value=mock_sync_client, ), patch( @@ -205,7 +199,7 @@ async def test_cyberark_rotate_secret(): current_value = {"value": initial_key_value} # Mock sync httpx client (for auth, ensure variable exists, sync reads) - # The _get_httpx_client returns an HTTPHandler with a .client property + # The get_httpx_client returns an HTTPHandler with a .client property mock_sync_client = MagicMock() # Auth response - note: the actual client is accessed via .client property mock_sync_client.client.post.return_value = create_mock_response( @@ -238,7 +232,7 @@ async def test_cyberark_rotate_secret(): with ( patch( - "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", return_value=mock_sync_client, ), patch( @@ -316,7 +310,7 @@ async def test_cyberark_rotate_secret_with_new_alias(): secrets_store = {} # Mock sync httpx client (for auth, ensure variable exists, sync reads) - # The _get_httpx_client returns an HTTPHandler with a .client property + # The get_httpx_client returns an HTTPHandler with a .client property mock_sync_client = MagicMock() # Auth response - note: the actual client is accessed via .client property mock_sync_client.client.post.return_value = create_mock_response( @@ -365,7 +359,7 @@ async def test_cyberark_rotate_secret_with_new_alias(): with ( patch( - "litellm.secret_managers.cyberark_secret_manager._get_httpx_client", + "litellm.secret_managers.cyberark_secret_manager.get_httpx_client", return_value=mock_sync_client, ), patch( diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index a0589136996..a888783c2ac 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -1009,9 +1009,7 @@ def test_async_http_handler(mock_async_client): concurrent_limit = 2 # Mock the transport creation to return a specific transport - with mock.patch.object( - AsyncHTTPHandler, "_create_async_transport" - ) as mock_create_transport: + with mock.patch.object(AsyncHTTPHandler, "create_async_transport") as mock_create_transport: mock_transport = mock.MagicMock() mock_create_transport.return_value = mock_transport diff --git a/tests/llm_translation/test_anthropic_completion.py b/tests/llm_translation/test_anthropic_completion.py index a79240d477d..92b0f97b0a4 100644 --- a/tests/llm_translation/test_anthropic_completion.py +++ b/tests/llm_translation/test_anthropic_completion.py @@ -436,7 +436,7 @@ def test_anthropic_tool_helper(cache_control_location): else: tool["cache_control"] = {"type": "ephemeral"} - tool, _ = AnthropicConfig()._map_tool_helper(tool=tool) + tool, _ = AnthropicConfig().map_tool_helper(tool=tool) assert tool["cache_control"] == {"type": "ephemeral"} @@ -581,7 +581,7 @@ def test_convert_tool_response_to_message_with_values(): ) ] - message = AnthropicConfig._convert_tool_response_to_message(tool_calls=tool_calls) + message = AnthropicConfig.convert_tool_response_to_message(tool_calls=tool_calls) assert message is not None assert message.content == '{"name": "John", "age": 30}' @@ -606,7 +606,7 @@ def test_convert_tool_response_to_message_without_values(): ) ] - message = AnthropicConfig._convert_tool_response_to_message(tool_calls=tool_calls) + message = AnthropicConfig.convert_tool_response_to_message(tool_calls=tool_calls) assert message is not None assert message.content == '{"name": "John", "age": 30}' @@ -618,14 +618,12 @@ def test_convert_tool_response_to_message_invalid_json(): ChatCompletionToolCallChunk( id="test_id", type="function", - function=ChatCompletionToolCallFunctionChunk( - name="json_tool_call", arguments="invalid json" - ), + function=ChatCompletionToolCallFunctionChunk(name="json_tool_call", arguments="invalid json"), index=0, ) ] - message = AnthropicConfig._convert_tool_response_to_message(tool_calls=tool_calls) + message = AnthropicConfig.convert_tool_response_to_message(tool_calls=tool_calls) assert message is not None assert message.content == "invalid json" @@ -642,15 +640,16 @@ def test_convert_tool_response_to_message_no_arguments(): ) ] - message = AnthropicConfig._convert_tool_response_to_message(tool_calls=tool_calls) + message = AnthropicConfig.convert_tool_response_to_message(tool_calls=tool_calls) assert message is None def test_anthropic_tool_with_image(): - from litellm.litellm_core_utils.prompt_templates.factory import prompt_factory import json + from litellm.litellm_core_utils.prompt_templates.factory import prompt_factory + b64_data = "iVBORw0KGgoAAAANSUhEu6U3//C9t/fKv5wDgpP1r5796XwC4zyH1D565bHGDqbY85AMb0nIQe+u3J390Xbtb9XgXxcK0/aqRXpdYcwgARbCN03FJk" image_url = f"data:image/png;base64,{b64_data}" messages = [ @@ -915,7 +914,7 @@ def test_map_stop_sequences(stop_input, expected_output, drop_params): """Test the _map_stop_sequences method of AnthropicConfig""" litellm.drop_params = drop_params config = AnthropicConfig() - result = config._map_stop_sequences(stop_input) + result = config.map_stop_sequences(stop_input) assert result == expected_output diff --git a/tests/llm_translation/test_azure_agents.py b/tests/llm_translation/test_azure_agents.py index e0741471582..b26e11d9842 100644 --- a/tests/llm_translation/test_azure_agents.py +++ b/tests/llm_translation/test_azure_agents.py @@ -174,19 +174,15 @@ def test_azure_ai_agents_config_get_agent_id(): config = AzureAIAgentsConfig() # Test with full model name - agent_id = config._get_agent_id("azure_ai/agents/asst_abc123", {}) + agent_id = config.get_agent_id("azure_ai/agents/asst_abc123", {}) assert agent_id == "asst_abc123" # Test with optional_params override - agent_id = config._get_agent_id( - "azure_ai/agents/asst_abc123", {"agent_id": "asst_override"} - ) + agent_id = config.get_agent_id("azure_ai/agents/asst_abc123", {"agent_id": "asst_override"}) assert agent_id == "asst_override" # Test with assistant_id in optional_params - agent_id = config._get_agent_id( - "azure_ai/agents/asst_abc123", {"assistant_id": "asst_assistant"} - ) + agent_id = config.get_agent_id("azure_ai/agents/asst_abc123", {"assistant_id": "asst_assistant"}) assert agent_id == "asst_assistant" diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 8496d5c63ec..a7836d2486e 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -1620,18 +1620,13 @@ def test_bedrock_nova_topk(top_k_param): captured_data = result return result - with patch( - "litellm.AmazonConverseConfig._transform_request", side_effect=mock_transform - ): + with patch("litellm.AmazonConverseConfig._transform_request", side_effect=mock_transform): litellm.completion(**data) # Assert that additionalRequestParameters exists and contains topK assert "additionalModelRequestFields" in captured_data assert "inferenceConfig" in captured_data["additionalModelRequestFields"] - assert ( - captured_data["additionalModelRequestFields"]["inferenceConfig"]["topK"] - == 10 - ) + assert captured_data["additionalModelRequestFields"]["inferenceConfig"]["topK"] == 10 def test_bedrock_cross_region_inference(monkeypatch): @@ -1712,10 +1707,11 @@ class TestBedrockEmbedding(BaseLLMEmbeddingTest): "inference_params": {}, } - transformed_request = ( - AmazonTitanMultimodalEmbeddingG1Config()._transform_request(**args) + transformed_request = AmazonTitanMultimodalEmbeddingG1Config().transform_request(**args) + assert ( + transformed_request["inputImage"] + == "iVBORw0KGgoAAAANSUhEUgAAAGQAAABkBAMAAACCzIhnAAAAG1BMVEURAAD///+ln5/h39/Dv79qX18uHx+If39MPz9oMSdmAAAACXBIWXMAAA7EAAAOxAGVKw4bAAABB0lEQVRYhe2SzWrEIBCAh2A0jxEs4j6GLDS9hqWmV5Flt0cJS+lRwv742DXpEjY1kOZW6HwHFZnPmVEBEARBEARB/jd0KYA/bcUYbPrRLh6amXHJ/K+ypMoyUaGthILzw0l+xI0jsO7ZcmCcm4ILd+QuVYgpHOmDmz6jBeJImdcUCmeBqQpuqRIbVmQsLCrAalrGpfoEqEogqbLTWuXCPCo+Ki1XGqgQ+jVVuhB8bOaHkvmYuzm/b0KYLWwoK58oFqi6XfxQ4Uz7d6WeKpna6ytUs5e8betMcqAv5YPC5EZB2Lm9FIn0/VP6R58+/GEY1X1egVoZ/3bt/EqF6malgSAIgiDIH+QL41409QMY0LMAAAAASUVORK5CYII=" ) - assert transformed_request["inputImage"] == "iVBORw0KGgoAAAANSUhEUgAAAGQAAABkBAMAAACCzIhnAAAAG1BMVEURAAD///+ln5/h39/Dv79qX18uHx+If39MPz9oMSdmAAAACXBIWXMAAA7EAAAOxAGVKw4bAAABB0lEQVRYhe2SzWrEIBCAh2A0jxEs4j6GLDS9hqWmV5Flt0cJS+lRwv742DXpEjY1kOZW6HwHFZnPmVEBEARBEARB/jd0KYA/bcUYbPrRLh6amXHJ/K+ypMoyUaGthILzw0l+xI0jsO7ZcmCcm4ILd+QuVYgpHOmDmz6jBeJImdcUCmeBqQpuqRIbVmQsLCrAalrGpfoEqEogqbLTWuXCPCo+Ki1XGqgQ+jVVuhB8bOaHkvmYuzm/b0KYLWwoK58oFqi6XfxQ4Uz7d6WeKpna6ytUs5e8betMcqAv5YPC5EZB2Lm9FIn0/VP6R58+/GEY1X1egVoZ/3bt/EqF6malgSAIgiDIH+QL41409QMY0LMAAAAASUVORK5CYII=" @pytest.mark.asyncio diff --git a/tests/llm_translation/test_bedrock_invoke_tests.py b/tests/llm_translation/test_bedrock_invoke_tests.py index 3f84723ed2f..7920a372e96 100644 --- a/tests/llm_translation/test_bedrock_invoke_tests.py +++ b/tests/llm_translation/test_bedrock_invoke_tests.py @@ -133,7 +133,7 @@ def test_nova_invoke_streaming_chunk_parsing(): "contentBlockIndex": 0, } } - result = decoder._chunk_parser(nova_text_chunk) + result = decoder.chunk_parser(nova_text_chunk) assert result.choices[0].delta.content == "Hello, how can I help?" assert result.choices[0].index == 0 assert not result.choices[0].finish_reason @@ -146,7 +146,7 @@ def test_nova_invoke_streaming_chunk_parsing(): "contentBlockIndex": 1, } } - result = decoder._chunk_parser(nova_tool_start_chunk) + result = decoder.chunk_parser(nova_tool_start_chunk) assert result.choices[0].delta.content == "" assert result.choices[0].index == 0 assert result.choices[0].delta.tool_calls is not None @@ -161,14 +161,11 @@ def test_nova_invoke_streaming_chunk_parsing(): "contentBlockIndex": 2, } } - result = decoder._chunk_parser(nova_tool_args_chunk) + result = decoder.chunk_parser(nova_tool_args_chunk) assert result.choices[0].delta.content == "" assert result.choices[0].index == 0 assert result.choices[0].delta.tool_calls is not None - assert ( - result.choices[0].delta.tool_calls[0].function.arguments - == '{"location": "New York"}' - ) + assert result.choices[0].delta.tool_calls[0].function.arguments == '{"location": "New York"}' # Test case 4: Stop reason in contentBlockDelta nova_stop_chunk = { @@ -176,6 +173,6 @@ def test_nova_invoke_streaming_chunk_parsing(): "stopReason": "tool_use", } } - result = decoder._chunk_parser(nova_stop_chunk) + result = decoder.chunk_parser(nova_stop_chunk) print(result) assert result.choices[0].finish_reason == "tool_calls" diff --git a/tests/llm_translation/test_bedrock_nova_embedding.py b/tests/llm_translation/test_bedrock_nova_embedding.py index cb04b7c1398..23be2e37f43 100644 --- a/tests/llm_translation/test_bedrock_nova_embedding.py +++ b/tests/llm_translation/test_bedrock_nova_embedding.py @@ -31,7 +31,7 @@ class TestNovaTransformationRequest: "truncation_mode": "END", } - request = config._transform_request( + request = config.transform_request( input="Hello, world!", inference_params=inference_params, async_invoke_route=False, @@ -61,7 +61,7 @@ class TestNovaTransformationRequest: "output_s3_uri": "s3://my-bucket/output/", } - request = config._transform_request( + request = config.transform_request( input="Long text content...", inference_params=inference_params, async_invoke_route=True, @@ -99,7 +99,7 @@ class TestNovaTransformationRequest: }, } - request = config._transform_request( + request = config.transform_request( input=image_data, inference_params=inference_params, async_invoke_route=False, @@ -127,7 +127,7 @@ class TestNovaTransformationRequest: }, } - request = config._transform_request( + request = config.transform_request( input="s3://my-bucket/video.mp4", inference_params=inference_params, async_invoke_route=False, @@ -155,7 +155,7 @@ class TestNovaTransformationRequest: }, } - request = config._transform_request( + request = config.transform_request( input="s3://my-bucket/audio.mp3", inference_params=inference_params, async_invoke_route=False, @@ -178,7 +178,7 @@ class TestNovaTransformationRequest: } with pytest.raises(ValueError, match="output_s3_uri is required"): - config._transform_request( + config.transform_request( input="Test text", inference_params=inference_params, async_invoke_route=True, @@ -190,7 +190,7 @@ class TestNovaTransformationRequest: """Test default embedding purpose is GENERIC_INDEX.""" config = AmazonNovaEmbeddingConfig() - request = config._transform_request( + request = config.transform_request( input="Test text", inference_params={}, async_invoke_route=False, @@ -203,7 +203,7 @@ class TestNovaTransformationRequest: """Test default embedding dimension is 3072.""" config = AmazonNovaEmbeddingConfig() - request = config._transform_request( + request = config.transform_request( input="Test text", inference_params={}, async_invoke_route=False, @@ -219,7 +219,7 @@ class TestNovaTransformationRequest: # Test with JPEG image data URL jpeg_data_url = "data:image/jpeg;base64,/9j/4AAQSkZJRgABAQAASABIAAD" - request = config._transform_request( + request = config.transform_request( input=jpeg_data_url, inference_params={"dimensions": 1024}, async_invoke_route=False, @@ -238,11 +238,9 @@ class TestNovaTransformationRequest: config = AmazonNovaEmbeddingConfig() # Test with PNG image data URL - png_data_url = ( - "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJ" - ) + png_data_url = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJ" - request = config._transform_request( + request = config.transform_request( input=png_data_url, inference_params={}, async_invoke_route=False, @@ -263,16 +261,14 @@ class TestNovaTransformationRequest: # Test with jpg (should be converted to jpeg) jpg_data_url = "data:image/jpg;base64,/9j/4AAQSkZJRg" - request = config._transform_request( + request = config.transform_request( input=jpg_data_url, inference_params={}, async_invoke_route=False, ) params = request["singleEmbeddingParams"] - assert ( - params["image"]["format"] == "jpeg" - ) # Should be converted from jpg to jpeg + assert params["image"]["format"] == "jpeg" # Should be converted from jpg to jpeg def test_data_url_video_parsing(self): """Test that data URL videos are properly parsed.""" @@ -280,7 +276,7 @@ class TestNovaTransformationRequest: video_data_url = "data:video/mp4;base64,AAAAIGZ0eXBpc29t" - request = config._transform_request( + request = config.transform_request( input=video_data_url, inference_params={}, async_invoke_route=False, @@ -297,7 +293,7 @@ class TestNovaTransformationRequest: audio_data_url = "data:audio/mp3;base64,SUQzBAAAAAAAI1RTU0UAAAA" - request = config._transform_request( + request = config.transform_request( input=audio_data_url, inference_params={}, async_invoke_route=False, @@ -327,9 +323,7 @@ class TestNovaTransformationResponse: } ] - result = config._transform_response( - response_list, model="amazon.nova-2-multimodal-embeddings-v1:0" - ) + result = config.transform_response(response_list, model="amazon.nova-2-multimodal-embeddings-v1:0") assert result.model == "amazon.nova-2-multimodal-embeddings-v1:0" assert len(result.data) == 1 @@ -361,9 +355,7 @@ class TestNovaTransformationResponse: }, ] - result = config._transform_response( - response_list, model="amazon.nova-2-multimodal-embeddings-v1:0" - ) + result = config.transform_response(response_list, model="amazon.nova-2-multimodal-embeddings-v1:0") assert len(result.data) == 2 assert result.data[0].embedding == [0.1, 0.2, 0.3] @@ -390,9 +382,7 @@ class TestNovaTransformationResponse: } ] - result = config._transform_response( - response_list, model="amazon.nova-2-multimodal-embeddings-v1:0" - ) + result = config.transform_response(response_list, model="amazon.nova-2-multimodal-embeddings-v1:0") assert len(result.data) == 2 assert result.data[0].embedding == [0.1, 0.2, 0.3] @@ -429,7 +419,7 @@ class TestNovaTransformationResponse: } ] - result = config._transform_response( + result = config.transform_response( response_list=response_list, model="amazon.nova-2-multimodal-embeddings-v1:0", batch_data=batch_data, @@ -467,7 +457,7 @@ class TestNovaTransformationResponse: } ] - result = config._transform_response( + result = config.transform_response( response_list=response_list, model="amazon.nova-2-multimodal-embeddings-v1:0", batch_data=batch_data, @@ -492,7 +482,7 @@ class TestNovaTransformationResponse: ] # Call without batch_data — should not break - result = config._transform_response( + result = config.transform_response( response_list=response_list, model="amazon.nova-2-multimodal-embeddings-v1:0", ) @@ -505,13 +495,9 @@ class TestNovaTransformationResponse: """Test async invoke response transformation.""" config = AmazonNovaEmbeddingConfig() - response = { - "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123" - } + response = {"invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123"} - result = config._transform_async_invoke_response( - response, model="amazon.nova-2-multimodal-embeddings-v1:0" - ) + result = config.transform_async_invoke_response(response, model="amazon.nova-2-multimodal-embeddings-v1:0") assert result.model == "amazon.nova-2-multimodal-embeddings-v1:0" assert len(result.data) == 1 diff --git a/tests/llm_translation/test_huggingface_chat_completion.py b/tests/llm_translation/test_huggingface_chat_completion.py index 90e6c2adb8d..4dff03f514f 100644 --- a/tests/llm_translation/test_huggingface_chat_completion.py +++ b/tests/llm_translation/test_huggingface_chat_completion.py @@ -148,20 +148,18 @@ PROVIDER_MAPPING_RESPONSE = { @pytest.fixture def mock_provider_mapping(): - with patch( - "litellm.llms.huggingface.chat.transformation._fetch_inference_provider_mapping" - ) as mock: + with patch("litellm.llms.huggingface.chat.transformation.fetch_inference_provider_mapping") as mock: mock.return_value = PROVIDER_MAPPING_RESPONSE yield mock @pytest.fixture(autouse=True) def clear_lru_cache(): - from litellm.llms.huggingface.common_utils import _fetch_inference_provider_mapping + from litellm.llms.huggingface.common_utils import fetch_inference_provider_mapping - _fetch_inference_provider_mapping.cache_clear() + fetch_inference_provider_mapping.cache_clear() yield - _fetch_inference_provider_mapping.cache_clear() + fetch_inference_provider_mapping.cache_clear() @pytest.fixture diff --git a/tests/llm_translation/test_lambda_ai.py b/tests/llm_translation/test_lambda_ai.py index e6f8b13d4ba..5c1954fb84a 100644 --- a/tests/llm_translation/test_lambda_ai.py +++ b/tests/llm_translation/test_lambda_ai.py @@ -23,7 +23,7 @@ def test_lambda_ai_get_openai_compatible_provider_info(): # Test with default values (no env vars set) with mock.patch.dict(os.environ, {}, clear=True): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://api.lambda.ai/v1" assert api_key is None @@ -35,7 +35,7 @@ def test_lambda_ai_get_openai_compatible_provider_info(): "LAMBDA_API_BASE": "https://custom.lambda.ai/v1", }, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://custom.lambda.ai/v1" assert api_key == "test-key" @@ -44,9 +44,7 @@ def test_lambda_ai_get_openai_compatible_provider_info(): os.environ, {"LAMBDA_API_KEY": "env-key", "LAMBDA_API_BASE": "https://env.lambda.ai/v1"}, ): - api_base, api_key = config._get_openai_compatible_provider_info( - "https://param.lambda.ai/v1", "param-key" - ) + api_base, api_key = config.get_openai_compatible_provider_info("https://param.lambda.ai/v1", "param-key") assert api_base == "https://param.lambda.ai/v1" assert api_key == "param-key" diff --git a/tests/llm_translation/test_prompt_factory.py b/tests/llm_translation/test_prompt_factory.py index 7b03736920b..d718b015d23 100644 --- a/tests/llm_translation/test_prompt_factory.py +++ b/tests/llm_translation/test_prompt_factory.py @@ -5,6 +5,7 @@ import pytest from typing import List +from unittest.mock import MagicMock, patch # from litellm.litellm_core_utils.prompt_templates.factory import prompt_factory import litellm @@ -26,10 +27,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_completion_messages, ) from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) from litellm.types.llms.openai import AllMessageValues -from unittest.mock import MagicMock, patch def test_llama_3_prompt(): @@ -562,9 +562,7 @@ def test_vertex_only_image_user_message(): }, ] - response = _gemini_convert_messages_with_history( - messages=messages, model="gemini-1.5-pro" - ) + response = gemini_convert_messages_with_history(messages=messages, model="gemini-1.5-pro") expected_response = [ { @@ -597,7 +595,7 @@ def test_no_messages_yields_user_text(): """ messages: List[AllMessageValues] = [] - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) expected_output = [{"role": "user", "parts": [{"text": " "}]}] diff --git a/tests/llm_translation/test_replicate.py b/tests/llm_translation/test_replicate.py index 129cb756224..436ed23c8f1 100644 --- a/tests/llm_translation/test_replicate.py +++ b/tests/llm_translation/test_replicate.py @@ -109,7 +109,7 @@ class TestReplicateStartingStatus: # Verify that GET was called 3 times (starting, processing, succeeded) assert mock_client.get.call_count == 3 - @patch("litellm.llms.replicate.chat.handler._get_httpx_client") + @patch("litellm.llms.replicate.chat.handler.get_httpx_client") def test_sync_completion_handles_starting_status(self, mock_get_client): """Test that sync completion polls correctly when status is 'starting'""" # Mock the sync HTTP client diff --git a/tests/llm_translation/test_v0.py b/tests/llm_translation/test_v0.py index e96022e1e22..7fcfa0cfc92 100644 --- a/tests/llm_translation/test_v0.py +++ b/tests/llm_translation/test_v0.py @@ -23,25 +23,19 @@ def test_v0_get_openai_compatible_provider_info(): # Test with default values (no env vars set) with mock.patch.dict(os.environ, {}, clear=True): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://api.v0.dev/v1" assert api_key is None # Test with environment variables - with mock.patch.dict( - os.environ, {"V0_API_KEY": "test-key", "V0_API_BASE": "https://custom.v0.ai/v1"} - ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + with mock.patch.dict(os.environ, {"V0_API_KEY": "test-key", "V0_API_BASE": "https://custom.v0.ai/v1"}): + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://custom.v0.ai/v1" assert api_key == "test-key" # Test with explicit parameters (should override env vars) - with mock.patch.dict( - os.environ, {"V0_API_KEY": "env-key", "V0_API_BASE": "https://env.v0.ai/v1"} - ): - api_base, api_key = config._get_openai_compatible_provider_info( - "https://param.v0.ai/v1", "param-key" - ) + with mock.patch.dict(os.environ, {"V0_API_KEY": "env-key", "V0_API_BASE": "https://env.v0.ai/v1"}): + api_base, api_key = config.get_openai_compatible_provider_info("https://param.v0.ai/v1", "param-key") assert api_base == "https://param.v0.ai/v1" assert api_key == "param-key" diff --git a/tests/llm_translation/test_xai.py b/tests/llm_translation/test_xai.py index 4f3346b5477..2c65922740e 100644 --- a/tests/llm_translation/test_xai.py +++ b/tests/llm_translation/test_xai.py @@ -1,35 +1,27 @@ import json import os from datetime import datetime -from unittest.mock import AsyncMock - - +from unittest.mock import AsyncMock, patch import httpx import pytest +from base_llm_unit_tests import BaseLLMChatTest, BaseReasoningLLMTests -from litellm import Choices, Message, ModelResponse, EmbeddingResponse, Usage -from litellm import completion -from unittest.mock import patch -from litellm.llms.xai.chat.transformation import XAIChatConfig, XAI_API_BASE -from base_llm_unit_tests import BaseReasoningLLMTests, BaseLLMChatTest +from litellm import Choices, EmbeddingResponse, Message, ModelResponse, Usage, completion +from litellm.llms.xai.chat.transformation import XAI_API_BASE, XAIChatConfig def test_xai_chat_config_get_openai_compatible_provider_info(): config = XAIChatConfig() # Test with default values - api_base, api_key = config._get_openai_compatible_provider_info( - api_base=None, api_key=None - ) + api_base, api_key = config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert api_base == XAI_API_BASE assert api_key == os.environ.get("XAI_API_KEY") # Test with custom API key custom_api_key = "test_api_key" - api_base, api_key = config._get_openai_compatible_provider_info( - api_base=None, api_key=custom_api_key - ) + api_base, api_key = config.get_openai_compatible_provider_info(api_base=None, api_key=custom_api_key) assert api_base == XAI_API_BASE assert api_key == custom_api_key @@ -38,7 +30,7 @@ def test_xai_chat_config_get_openai_compatible_provider_info(): "os.environ", {"XAI_API_BASE": "https://env.x.ai/v1", "XAI_API_KEY": "env_api_key"}, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://env.x.ai/v1" assert api_key == "env_api_key" diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py index 02ea8a54f50..a592935c892 100644 --- a/tests/local_testing/test_amazing_vertex_completion.py +++ b/tests/local_testing/test_amazing_vertex_completion.py @@ -22,7 +22,7 @@ from litellm import ( image_generation, ) from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) from litellm.llms.vertex_ai.vertex_llm_base import VertexBase @@ -1190,15 +1190,12 @@ def test_tool_name_conversion(): # Now the assistant can reply with the result of the tool call. ] - translated_messages = _gemini_convert_messages_with_history(messages=messages) + translated_messages = gemini_convert_messages_with_history(messages=messages) print(f"\n\ntranslated_messages: {translated_messages}\ntranslated_messages") # assert that the last tool response has the corresponding tool name - assert ( - translated_messages[-1]["parts"][0]["function_response"]["name"] - == "get_weather" - ) + assert translated_messages[-1]["parts"][0]["function_response"]["name"] == "get_weather" def test_prompt_factory(): @@ -1237,7 +1234,7 @@ def test_prompt_factory(): # Now the assistant can reply with the result of the tool call. ] - translated_messages = _gemini_convert_messages_with_history(messages=messages) + translated_messages = gemini_convert_messages_with_history(messages=messages) print(f"\n\ntranslated_messages: {translated_messages}\ntranslated_messages") @@ -1247,23 +1244,19 @@ def test_prompt_factory_nested(): {"role": "user", "content": [{"type": "text", "text": "hi"}]}, { "role": "assistant", - "content": [ - {"type": "text", "text": "Hi! 👋 \n\nHow can I help you today? 😊 \n"} - ], + "content": [{"type": "text", "text": "Hi! 👋 \n\nHow can I help you today? 😊 \n"}], }, {"role": "user", "content": [{"type": "text", "text": "hi 2nd time"}]}, ] - translated_messages = _gemini_convert_messages_with_history(messages=messages) + translated_messages = gemini_convert_messages_with_history(messages=messages) print(f"\n\ntranslated_messages: {translated_messages}\ntranslated_messages") for message in translated_messages: assert len(message["parts"]) == 1 assert "text" in message["parts"][0], "Missing 'text' from 'parts'" - assert isinstance( - message["parts"][0]["text"], str - ), "'text' value not a string." + assert isinstance(message["parts"][0]["text"], str), "'text' value not a string." @pytest.mark.asyncio @@ -1942,7 +1935,7 @@ def test_gemini_function_call_parameter_in_messages(): def test_gemini_function_call_parameter_in_messages_2(): litellm.set_verbose = True from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -1962,7 +1955,7 @@ def test_gemini_function_call_parameter_in_messages_2(): }, ] - returned_contents = _gemini_convert_messages_with_history(messages=messages) + returned_contents = gemini_convert_messages_with_history(messages=messages) print(f"returned_contents: {returned_contents}") assert returned_contents == [ diff --git a/tests/local_testing/test_embedding.py b/tests/local_testing/test_embedding.py index 6f963f349a8..83ebfa102a3 100644 --- a/tests/local_testing/test_embedding.py +++ b/tests/local_testing/test_embedding.py @@ -1184,9 +1184,7 @@ def test_encoding_format_explicit_value_preserved(): When user provides encoding_format='float' or 'base64', it should be sent as-is to the OpenAI SDK. """ - with patch( - "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" - ) as mock_get_client: + with patch("litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client") as mock_get_client: # Create a mock client instance mock_client_instance = MagicMock() mock_get_client.return_value = mock_client_instance diff --git a/tests/local_testing/test_ollama.py b/tests/local_testing/test_ollama.py index b76b64f8ce7..27d9a59c69e 100644 --- a/tests/local_testing/test_ollama.py +++ b/tests/local_testing/test_ollama.py @@ -227,14 +227,12 @@ async def test_async_ollama_ssl_verify(stream): except Exception as e: print(e) - client: AsyncHTTPHandler = litellm.in_memory_llm_clients_cache.get_cache( - "async_httpx_clientssl_verify_Falseollama" - ) + client: AsyncHTTPHandler = litellm.in_memory_llm_clients_cache.get_cache("async_httpx_clientssl_verify_Falseollama") # check client print("type of transport in client=", type(client.client._transport)) print("vars in transport in client=", vars(client.client._transport)) - litellm_created_session = client.client._transport._get_valid_client_session() + litellm_created_session = client.client._transport.get_valid_client_session() print("litellm_created_session=", litellm_created_session) # check session ssl print("litellm_created_session ssl=", litellm_created_session.connector._ssl) diff --git a/tests/router_unit_tests/test_router_endpoints.py b/tests/router_unit_tests/test_router_endpoints.py index df2f352e6ea..66bd294268a 100644 --- a/tests/router_unit_tests/test_router_endpoints.py +++ b/tests/router_unit_tests/test_router_endpoints.py @@ -369,7 +369,7 @@ async def test_moderation_endpoint_with_api_base(): # Mock the OpenAI client to verify api_base is passed with patch( - "litellm.main.openai_chat_completions._get_openai_client" + "litellm.main.openai_chat_completions.get_openai_client" ) as mock_get_client: mock_client = AsyncMock() mock_response = MagicMock() @@ -392,7 +392,7 @@ async def test_moderation_endpoint_with_api_base(): model="openai/omni-moderation-latest", input="hello this is a test" ) - # Verify that _get_openai_client was called with the custom api_base + # Verify that get_openai_client was called with the custom api_base mock_get_client.assert_called() call_kwargs = mock_get_client.call_args.kwargs assert ( diff --git a/tests/test_litellm/conftest.py b/tests/test_litellm/conftest.py index f83c1e76b3a..bd38d3b02e8 100644 --- a/tests/test_litellm/conftest.py +++ b/tests/test_litellm/conftest.py @@ -249,7 +249,7 @@ def isolate_litellm_state(): original_state["model_fallbacks"] = litellm.model_fallbacks # Store transport/network globals — many tests set these without restoring, - # causing subsequent tests to get None from _create_async_transport() + # causing subsequent tests to get None from create_async_transport() for _attr in ("disable_aiohttp_transport", "force_ipv4"): if hasattr(litellm, _attr): original_state[_attr] = getattr(litellm, _attr) diff --git a/tests/unit/batches/test_main.py b/tests/unit/batches/test_main.py index 26dc4083b0b..45a684c007b 100644 --- a/tests/unit/batches/test_main.py +++ b/tests/unit/batches/test_main.py @@ -117,7 +117,7 @@ def test_create__openai_dispatch_and_payload(seams): # DISPATCH + RESULT assert result is seams.openai.create_batch.return_value _assert_only(seams.openai.create_batch, seams, "create_batch") - seams.bedrock_arn._handle_async_invoke_status.assert_not_called() + seams.bedrock_arn.handle_async_invoke_status.assert_not_called() # PAYLOAD - request object built from the call, sync flag off. kw = seams.openai.create_batch.call_args.kwargs @@ -265,8 +265,8 @@ def test_retrieve__bedrock_async_invoke_arn(seams): arn = "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123" result = bm.retrieve_batch(batch_id=arn, custom_llm_provider="bedrock") - seams.bedrock_arn._handle_async_invoke_status.assert_called_once() - assert result is seams.bedrock_arn._handle_async_invoke_status.return_value + seams.bedrock_arn.handle_async_invoke_status.assert_called_once() + assert result is seams.bedrock_arn.handle_async_invoke_status.return_value # provider instances untouched. for m in _all_seam_methods(seams, "retrieve_batch"): m.assert_not_called() @@ -276,9 +276,9 @@ def test_retrieve__bedrock_model_invocation_job_arn(seams): arn = "arn:aws:bedrock:us-east-1:123456789012:model-invocation-job/xyz789" result = bm.retrieve_batch(batch_id=arn, custom_llm_provider="bedrock") - seams.bedrock_arn._handle_model_invocation_job_status.assert_called_once() - assert result is seams.bedrock_arn._handle_model_invocation_job_status.return_value - seams.bedrock_arn._handle_async_invoke_status.assert_not_called() + seams.bedrock_arn.handle_model_invocation_job_status.assert_called_once() + assert result is seams.bedrock_arn.handle_model_invocation_job_status.return_value + seams.bedrock_arn.handle_async_invoke_status.assert_not_called() def test_retrieve__unsupported_provider_raises_badrequest(seams): diff --git a/tests/unit/caching/test_gcs_cache.py b/tests/unit/caching/test_gcs_cache.py index 4dba0e76a57..fe5df6526be 100644 --- a/tests/unit/caching/test_gcs_cache.py +++ b/tests/unit/caching/test_gcs_cache.py @@ -14,12 +14,15 @@ def mock_gcs_dependencies(): mock_async_client = AsyncMock() with ( - patch.object(import_module("litellm.caching.gcs_cache"), "_get_httpx_client", return_value=mock_sync_client - ), - patch.object(import_module("litellm.caching.gcs_cache"), "get_async_httpx_client", + patch.object(import_module("litellm.caching.gcs_cache"), "get_httpx_client", return_value=mock_sync_client), + patch.object( + import_module("litellm.caching.gcs_cache"), + "get_async_httpx_client", return_value=mock_async_client, ), - patch.object(import_module("litellm.caching.gcs_cache").GCSBucketBase, "sync_construct_request_headers", + patch.object( + import_module("litellm.caching.gcs_cache").GCSBucketBase, + "sync_construct_request_headers", return_value={}, ), ): diff --git a/tests/unit/caching/test_llm_caching_handler.py b/tests/unit/caching/test_llm_caching_handler.py index b6c90ba288e..0c0fe8ca3c5 100644 --- a/tests/unit/caching/test_llm_caching_handler.py +++ b/tests/unit/caching/test_llm_caching_handler.py @@ -535,7 +535,7 @@ async def test_the_pool_is_released_once_the_stream_it_carried_ends(monkeypatch, def is_released() -> bool: return pool.connections == [] else: - session = transport._get_valid_client_session() + session = transport.get_valid_client_session() def is_released() -> bool: return session.closed diff --git a/tests/unit/caching/test_qdrant_semantic_cache.py b/tests/unit/caching/test_qdrant_semantic_cache.py index 4f18fb1bca6..b67cbd90b32 100644 --- a/tests/unit/caching/test_qdrant_semantic_cache.py +++ b/tests/unit/caching/test_qdrant_semantic_cache.py @@ -13,12 +13,9 @@ def test_qdrant_semantic_cache_initialization(monkeypatch): """ # Mock the httpx clients and API calls with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -75,12 +72,9 @@ def test_qdrant_semantic_cache_get_cache_hit(): Verifies that cached results are properly retrieved and parsed. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -155,12 +149,9 @@ def test_qdrant_semantic_cache_rejects_unscoped_cache_hit(): safely migrated to a generated LiteLLM cache key. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - mock_response = MagicMock() mock_response.status_code = 200 mock_response.json.return_value = {"result": {"exists": True}} @@ -316,12 +307,9 @@ def test_qdrant_semantic_cache_get_cache_miss(): Verifies that None is returned when no similar cached results are found. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -370,14 +358,9 @@ async def test_qdrant_semantic_cache_async_get_cache_hit(): Verifies that cached results are properly retrieved and parsed asynchronously. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, - patch( - "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" - ) as mock_async_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_async_client, ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -461,14 +444,9 @@ async def test_qdrant_semantic_cache_async_get_cache_miss(): Verifies that None is returned when no similar cached results are found. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, - patch( - "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" - ) as mock_async_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_async_client, ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -523,12 +501,9 @@ def test_qdrant_semantic_cache_set_cache(): Verifies that responses are properly stored in the cache. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -588,14 +563,9 @@ async def test_qdrant_semantic_cache_async_set_cache(): Verifies that responses are properly stored in the cache asynchronously. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, - patch( - "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" - ) as mock_async_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_async_client, ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -661,12 +631,9 @@ def test_qdrant_semantic_cache_custom_vector_size(): creation payload instead of the default 1536. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection does NOT exist (so it will be created) mock_exists_response = MagicMock() mock_exists_response.status_code = 200 @@ -722,12 +689,9 @@ def test_qdrant_semantic_cache_default_vector_size(): is not provided, and stores it as self.vector_size. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection exists check mock_response = MagicMock() mock_response.status_code = 200 @@ -758,12 +722,9 @@ def test_qdrant_semantic_cache_large_vector_size(): for models like Stella, bge-en-icl, etc. """ with ( - patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_sync_client, + patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_sync_client, patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client"), ): - # Mock the collection does NOT exist (so it will be created) mock_exists_response = MagicMock() mock_exists_response.status_code = 200 diff --git a/tests/unit/containers/test_container_integration.py b/tests/unit/containers/test_container_integration.py index 6c3a876fc45..22fd158983f 100644 --- a/tests/unit/containers/test_container_integration.py +++ b/tests/unit/containers/test_container_integration.py @@ -58,9 +58,7 @@ class TestContainerIntegration: mock_client.post.return_value = mock_response mock_http_handler.return_value = mock_client - with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client" - ) as mock_get_client: + with patch("litellm.llms.custom_httpx.llm_http_handler.get_httpx_client") as mock_get_client: mock_get_client.return_value = mock_client # Execute @@ -114,15 +112,11 @@ class TestContainerIntegration: mock_client.get.return_value = mock_response mock_http_handler.return_value = mock_client - with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client" - ) as mock_get_client: + with patch("litellm.llms.custom_httpx.llm_http_handler.get_httpx_client") as mock_get_client: mock_get_client.return_value = mock_client # Execute - response = list_containers( - limit=10, order="desc", custom_llm_provider="openai" - ) + response = list_containers(limit=10, order="desc", custom_llm_provider="openai") # Verify assert isinstance(response, ContainerListResponse) @@ -154,15 +148,11 @@ class TestContainerIntegration: mock_client.get.return_value = mock_response mock_http_handler.return_value = mock_client - with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client" - ) as mock_get_client: + with patch("litellm.llms.custom_httpx.llm_http_handler.get_httpx_client") as mock_get_client: mock_get_client.return_value = mock_client # Execute - response = retrieve_container( - container_id=container_id, custom_llm_provider="openai" - ) + response = retrieve_container(container_id=container_id, custom_llm_provider="openai") # Verify assert isinstance(response, ContainerObject) @@ -188,15 +178,11 @@ class TestContainerIntegration: mock_client.delete.return_value = mock_response mock_http_handler.return_value = mock_client - with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client" - ) as mock_get_client: + with patch("litellm.llms.custom_httpx.llm_http_handler.get_httpx_client") as mock_get_client: mock_get_client.return_value = mock_client # Execute - response = delete_container( - container_id=container_id, custom_llm_provider="openai" - ) + response = delete_container(container_id=container_id, custom_llm_provider="openai") # Verify assert isinstance(response, DeleteContainerResult) diff --git a/tests/unit/google_genai/test_google_genai_adapter.py b/tests/unit/google_genai/test_google_genai_adapter.py index 38292f471c8..5e521506699 100644 --- a/tests/unit/google_genai/test_google_genai_adapter.py +++ b/tests/unit/google_genai/test_google_genai_adapter.py @@ -1118,15 +1118,15 @@ def test_validate_environment_sets_x_goog_api_key(): def test_get_gemini_url_excludes_api_key(): """ - Verify that _get_gemini_url never embeds the API key in the URL. + Verify that get_gemini_url never embeds the API key in the URL. API keys in URLs leak through httpx error tracebacks. The key must be sent via the x-goog-api-key header instead. """ - from litellm.llms.vertex_ai.common_utils import _get_gemini_url + from litellm.llms.vertex_ai.common_utils import get_gemini_url for mode in ("chat", "embedding", "batch_embedding", "count_tokens"): - url, _ = _get_gemini_url( + url, _ = get_gemini_url( mode=mode, model="gemini-2.5-flash", stream=False, @@ -1134,7 +1134,7 @@ def test_get_gemini_url_excludes_api_key(): assert "key=" not in url, f"API key found in URL for mode={mode}: {url}" # Streaming chat should only have ?alt=sse - url, _ = _get_gemini_url(mode="chat", model="gemini-2.5-flash", stream=True) + url, _ = get_gemini_url(mode="chat", model="gemini-2.5-flash", stream=True) assert "key=" not in url, f"API key found in streaming URL: {url}" assert "alt=sse" in url, f"Missing alt=sse in streaming URL: {url}" diff --git a/tests/unit/google_genai/test_google_genai_transformation.py b/tests/unit/google_genai/test_google_genai_transformation.py index f0d0fc6126d..a95c8e5666f 100644 --- a/tests/unit/google_genai/test_google_genai_transformation.py +++ b/tests/unit/google_genai/test_google_genai_transformation.py @@ -381,7 +381,7 @@ def test_transform_generate_content_request_normalizes_response_schema_2_5(): def test_transform_generate_content_request_flattens_response_schema_1_5(): """For Gemini 1.5, ``responseSchema`` is kept but flattened via - ``_build_vertex_schema`` so ``$defs``/``$ref`` are unpacked.""" + ``build_vertex_schema`` so ``$defs``/``$ref`` are unpacked.""" config = GoogleGenAIConfig() schema = { diff --git a/tests/unit/integrations/gcs_bucket/test_gcs_bucket_base.py b/tests/unit/integrations/gcs_bucket/test_gcs_bucket_base.py index fb53994089b..781b56cae14 100644 --- a/tests/unit/integrations/gcs_bucket/test_gcs_bucket_base.py +++ b/tests/unit/integrations/gcs_bucket/test_gcs_bucket_base.py @@ -23,7 +23,7 @@ class TestGCSBucketBase: with ( patch( - "litellm.vertex_chat_completion._ensure_access_token" + "litellm.vertex_chat_completion.ensure_access_token" ) as mock_ensure_token, patch( "litellm.vertex_chat_completion._get_token_and_url" @@ -66,7 +66,7 @@ class TestGCSBucketBase: with ( patch( - "litellm.vertex_chat_completion._ensure_access_token" + "litellm.vertex_chat_completion.ensure_access_token" ) as mock_ensure_token, patch( "litellm.vertex_chat_completion._get_token_and_url" diff --git a/tests/unit/integrations/gcs_pubsub/test_pub_sub.py b/tests/unit/integrations/gcs_pubsub/test_pub_sub.py index 3c7f577d1d8..2196116c383 100644 --- a/tests/unit/integrations/gcs_pubsub/test_pub_sub.py +++ b/tests/unit/integrations/gcs_pubsub/test_pub_sub.py @@ -35,7 +35,7 @@ async def test_construct_request_headers_project_id_from_env(monkeypatch): mock_token = "mock-token" with patch( - "litellm.vertex_chat_completion._ensure_access_token_async" + "litellm.vertex_chat_completion.ensure_access_token_async" ) as mock_ensure_token: mock_ensure_token.return_value = (mock_auth_header, test_project_id) @@ -53,7 +53,7 @@ async def test_construct_request_headers_project_id_from_env(monkeypatch): "Content-Type": "application/json", } - # Verify _ensure_access_token_async was called with correct project_id + # Verify ensure_access_token_async was called with correct project_id mock_ensure_token.assert_called_once_with( credentials="test-path.json", project_id=test_project_id, diff --git a/tests/unit/integrations/langfuse/test_langfuse_prompt_management.py b/tests/unit/integrations/langfuse/test_langfuse_prompt_management.py index b487fe88b06..7e53055c9ac 100644 --- a/tests/unit/integrations/langfuse/test_langfuse_prompt_management.py +++ b/tests/unit/integrations/langfuse/test_langfuse_prompt_management.py @@ -68,9 +68,9 @@ class TestLangfusePromptManagement: def test_langfuse_client_init_passes_dedicated_httpx_client(self): import httpx - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client - shared_client = _get_httpx_client().client + shared_client = get_httpx_client().client built = MagicMock() with ( patch( diff --git a/tests/unit/integrations/langfuse/test_langfuse_sdk.py b/tests/unit/integrations/langfuse/test_langfuse_sdk.py index 1669b9233b0..5f01b35a338 100644 --- a/tests/unit/integrations/langfuse/test_langfuse_sdk.py +++ b/tests/unit/integrations/langfuse/test_langfuse_sdk.py @@ -1543,12 +1543,12 @@ def test_exporter_reports_failure_when_no_span_of_the_batch_can_be_encoded(monke def test_built_exporter_uses_the_shared_litellm_handler_and_langfuse_headers(monkeypatch): """No private requests session or TLS adapter: the channel is the same handler the rest of litellm uses.""" - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client monkeypatch.delenv("LANGFUSE_TIMEOUT", raising=False) monkeypatch.delenv("LANGFUSE_MAX_RETRIES", raising=False) default = _build_span_exporter(public_key="pk", secret_key="sk", base_url="https://lf.internal.example") - assert default.handler is _get_httpx_client() + assert default.handler is get_httpx_client() assert default.timeout == 20 assert len(default.delays) == 3 diff --git a/tests/unit/integrations/otel/test_otel_v2_presets.py b/tests/unit/integrations/otel/test_otel_v2_presets.py index a060cdf3648..64366626784 100644 --- a/tests/unit/integrations/otel/test_otel_v2_presets.py +++ b/tests/unit/integrations/otel/test_otel_v2_presets.py @@ -123,7 +123,7 @@ def test_agentops_exporter_tolerates_fetch_failure(monkeypatch): def test_fetch_jwt_uses_owned_client_not_shared_pool(monkeypatch): """The fetch owns a short-lived client and closes it, rather than closing - the process-wide cached ``_get_httpx_client`` pool shared by other callers.""" + the process-wide cached ``get_httpx_client`` pool shared by other callers.""" closed = {"n": 0} class _FakeResponse: @@ -146,7 +146,7 @@ def test_fetch_jwt_uses_owned_client_not_shared_pool(monkeypatch): return _FakeResponse() monkeypatch.setattr(httpx, "Client", _FakeClient) - assert not hasattr(agentops_mod, "_get_httpx_client") + assert not hasattr(agentops_mod, "get_httpx_client") result = _fetch_agentops_jwt("api-key") assert result == {"token": "jwt-123"} diff --git a/tests/unit/integrations/test_langfuse.py b/tests/unit/integrations/test_langfuse.py index 3e6e130cac5..bd55216e4bf 100644 --- a/tests/unit/integrations/test_langfuse.py +++ b/tests/unit/integrations/test_langfuse.py @@ -1399,13 +1399,12 @@ def test_langfuse_rest_client_survives_httpx_cache_eviction(monkeypatch): import weakref from litellm.caching.llm_caching_handler import LLMClientCache - - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + from litellm.llms.custom_httpx.http_handler import get_httpx_client monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) logger = _build_langfuse_logger(monkeypatch) - cached_handler = _get_httpx_client() + cached_handler = get_httpx_client() handler_ref = weakref.ref(cached_handler) assert logger.langfuse_client is cached_handler.client diff --git a/tests/unit/integrations/test_s3_v2.py b/tests/unit/integrations/test_s3_v2.py index 174c9ed2a06..8e9a8982cd9 100644 --- a/tests/unit/integrations/test_s3_v2.py +++ b/tests/unit/integrations/test_s3_v2.py @@ -149,7 +149,7 @@ class TestS3V2UnitTests: mock_sync_client.put.return_value = mock_response with patch( - "litellm.integrations.s3_v2._get_httpx_client", + "litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client, ): s3_logger_sync.upload_data_to_s3(test_element) @@ -290,7 +290,7 @@ class TestS3V2UnitTests: mock_sync_client.put.return_value = mock_response with patch( - "litellm.integrations.s3_v2._get_httpx_client", + "litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client, ): s3_logger_sync_virtual.upload_data_to_s3(test_element) @@ -781,7 +781,7 @@ def test_sync_upload_retries_403_with_fresh_signature(rotating_profile: str, mon handler.client = httpx.Client(transport=httpx.MockTransport(respond)) with ( patch( # test-quality-ok: sync upload builds its HTTPHandler per call, there is no injection seam for it - "litellm.integrations.s3_v2._get_httpx_client", return_value=handler + "litellm.integrations.s3_v2.get_httpx_client", return_value=handler ), patch("time.sleep") as mock_sleep, ): @@ -827,7 +827,7 @@ def test_sync_upload_retries_on_s3_503(): mock_sync_client.put = MagicMock(side_effect=[response_503, response_200]) with patch( - "litellm.integrations.s3_v2._get_httpx_client", + "litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client, ): with patch("time.sleep") as mock_sleep: @@ -1889,7 +1889,7 @@ def test_sync_upload_sets_content_md5_header(monkeypatch): mock_sync_client.put.return_value = response with patch( - "litellm.integrations.s3_v2._get_httpx_client", + "litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client, ): logger.upload_data_to_s3(test_element) @@ -2020,7 +2020,7 @@ def test_sync_upload_sets_sse_kms_key_id_header_when_configured(): mock_sync_client.put.return_value = response with patch( - "litellm.integrations.s3_v2._get_httpx_client", + "litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client, ): logger.upload_data_to_s3(test_element) @@ -2297,7 +2297,7 @@ def test_sync_upload_signs_object_key_with_space_the_way_s3_does(): mock_sync_client = MagicMock() mock_sync_client.put.return_value = response - with patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client): + with patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client): logger.upload_data_to_s3(_element_with_space()) call = mock_sync_client.put.call_args @@ -2395,7 +2395,7 @@ def test_sync_upload_percent_encodes_reserved_characters_in_object_key(s3_object mock_sync_client = MagicMock() mock_sync_client.put.return_value = response - with patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client): + with patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client): logger.upload_data_to_s3(_element_for(s3_object_key)) call = mock_sync_client.put.call_args @@ -3749,7 +3749,7 @@ def test_sync_upload_404_is_single_attempt_without_sleep() -> None: sync_client: Final = _SyncRecordingClient(_coded_failure_response(404, "NoSuchKey")) with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=sync_client), patch("time.sleep") as mock_sleep, ): logger.upload_data_to_s3(_element({"id": "sync-404"}, "sync-404")) @@ -3779,7 +3779,7 @@ def test_sync_upload_retry_set_matches_base(status: int, expected_puts: int, exp sync_client: Final = _SyncRecordingClient(_coded_failure_response(status, "SlowDown")) with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=sync_client), patch("time.sleep") as mock_sleep, ): logger.upload_data_to_s3(_element({"id": "sync"}, "sync")) @@ -4135,7 +4135,7 @@ def test_sync_terminal_code_is_retried_like_base_when_the_drop_flag_is_off() -> logger.handle_callback_failure = failures with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client), patch("time.sleep"), ): logger.upload_data_to_s3(_element({"id": "x"}, "x")) @@ -4152,7 +4152,7 @@ def test_sync_terminal_code_is_retried_like_base_when_the_drop_flag_is_off() -> mock_sync_client.put = MagicMock(side_effect=[_coded_failure_response(403, "InvalidRequest"), _ok_response()]) with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client), patch("time.sleep"), ): dropping.upload_data_to_s3(_element({"id": "x"}, "x")) @@ -4174,7 +4174,7 @@ def test_sync_retry_lines_stay_at_warning_level(caplog) -> None: with ( caplog.at_level("WARNING"), - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client), patch("time.sleep"), ): logger.upload_data_to_s3(_element({"id": "x"}, "x")) @@ -4252,7 +4252,7 @@ def test_init_bypassed_sync_logger_retries_a_503_and_reports_a_404() -> None: mock_sync_client.put = MagicMock(side_effect=[_transient_failure_response(503), _ok_response()]) with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client), patch("time.sleep"), ): logger.upload_data_to_s3(_element({"id": "x"}, "x")) @@ -4263,7 +4263,7 @@ def test_init_bypassed_sync_logger_retries_a_503_and_reports_a_404() -> None: mock_sync_client.put = MagicMock(return_value=_coded_failure_response(404, None)) with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client), patch("time.sleep"), ): logger.upload_data_to_s3(_element({"id": "y"}, "y")) @@ -4682,7 +4682,7 @@ def test_sync_upload_retries_access_denied_403(caplog): mock_sync_client.put = MagicMock(return_value=_coded_failure_response(403, "AccessDenied")) with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client), patch("time.sleep") as mock_sleep, ): logger.upload_data_to_s3(test_element) @@ -4711,7 +4711,7 @@ def test_sync_upload_drops_terminal_object_once_and_logs_it_only_when_opted_in(c mock_sync_client.put = MagicMock(return_value=_coded_failure_response(400, "EntityTooLarge")) with ( - patch("litellm.integrations.s3_v2._get_httpx_client", return_value=mock_sync_client), + patch("litellm.integrations.s3_v2.get_httpx_client", return_value=mock_sync_client), patch("time.sleep") as mock_sleep, ): logger.upload_data_to_s3(test_element) diff --git a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py index a98b0d0115f..f683c4acbed 100644 --- a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py @@ -1239,7 +1239,7 @@ def test_handle_error_marks_only_a_status_code_it_never_received(): handler = BaseLLMHTTPHandler() with pytest.raises(litellm.llms.base_llm.chat.transformation.BaseLLMException) as transport: - raise handler._handle_error(e=httpx.ConnectError("Connection refused"), provider_config=None) + raise handler.handle_error(e=httpx.ConnectError("Connection refused"), provider_config=None) assert transport.value.status_code == 500 assert transport.value.status_code_is_synthesized is True @@ -1250,7 +1250,7 @@ def test_handle_error_marks_only_a_status_code_it_never_received(): response=httpx.Response(status_code=500, request=request, text="upstream exploded"), ) with pytest.raises(litellm.llms.base_llm.chat.transformation.BaseLLMException) as received: - raise handler._handle_error(e=upstream, provider_config=None) + raise handler.handle_error(e=upstream, provider_config=None) assert received.value.status_code == 500 assert received.value.status_code_is_synthesized is False diff --git a/tests/unit/litellm_core_utils/test_fallback_generalizations.py b/tests/unit/litellm_core_utils/test_fallback_generalizations.py index f70b52a7026..e10751b2c71 100644 --- a/tests/unit/litellm_core_utils/test_fallback_generalizations.py +++ b/tests/unit/litellm_core_utils/test_fallback_generalizations.py @@ -741,8 +741,8 @@ def test_shipped_adaptive_rule_gates_on_version_not_pricing(shipped_cost_map): non_adaptive = "us.anthropic.claude-opus-4-20250514" assert adaptive not in litellm.model_cost assert non_adaptive not in litellm.model_cost - assert AnthropicModelInfo._is_adaptive_thinking_model(adaptive, "anthropic") is True - assert AnthropicModelInfo._is_adaptive_thinking_model(non_adaptive, "anthropic") is False + assert AnthropicModelInfo.is_adaptive_thinking_model(adaptive, "anthropic") is True + assert AnthropicModelInfo.is_adaptive_thinking_model(non_adaptive, "anthropic") is False def test_shipped_rules_resolve_unmapped_future_bedrock_claude_with_both_flags(shipped_cost_map): diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index 51cf0a057b7..95a22bcbd67 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -462,7 +462,7 @@ async def test_realtime_health_check_uses_model_level_vertex_params(): fake_vertex_base = MagicMock() fake_vertex_base.get_vertex_region = MagicMock(return_value="us-central1") - fake_vertex_base._ensure_access_token_async = AsyncMock(return_value=("model-level-token", "model-level-project")) + fake_vertex_base.ensure_access_token_async = AsyncMock(return_value=("model-level-token", "model-level-project")) connect_calls = [] with ( @@ -491,7 +491,7 @@ async def test_realtime_health_check_uses_model_level_vertex_params(): fake_vertex_base.get_vertex_region.assert_called_once_with( vertex_region="us-central1", model="gemini-live-2.5-flash-native-audio" ) - fake_vertex_base._ensure_access_token_async.assert_called_once_with( + fake_vertex_base.ensure_access_token_async.assert_called_once_with( credentials='{"type":"service_account"}', project_id="model-level-project", custom_llm_provider="vertex_ai", diff --git a/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py index 332153b4c7d..92881ea82aa 100644 --- a/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/unit/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -608,7 +608,7 @@ def test_map_tool_helper(): tool = {"type": "web_search_20250305", "name": "web_search", "max_uses": 5} - result, _ = config._map_tool_helper(tool) + result, _ = config.map_tool_helper(tool) assert result is not None assert result["name"] == "web_search" assert result["max_uses"] == 5 @@ -1424,7 +1424,7 @@ def test_tool_search_regex_mapping(): tool = {"type": "tool_search_tool_regex_20251119", "name": "tool_search_tool_regex"} - mapped_tool, mcp_server = config._map_tool_helper(tool) + mapped_tool, mcp_server = config.map_tool_helper(tool) assert mapped_tool is not None assert mapped_tool["type"] == "tool_search_tool_regex_20251119" @@ -1438,7 +1438,7 @@ def test_tool_search_bm25_mapping(): tool = {"type": "tool_search_tool_bm25_20251119", "name": "tool_search_tool_bm25"} - mapped_tool, mcp_server = config._map_tool_helper(tool) + mapped_tool, mcp_server = config.map_tool_helper(tool) assert mapped_tool is not None assert mapped_tool["type"] == "tool_search_tool_bm25_20251119" @@ -1560,7 +1560,7 @@ def test_defer_loading_preserved_in_transformation(): "defer_loading": True, } - mapped_tool, mcp_server = config._map_tool_helper(tool) + mapped_tool, mcp_server = config.map_tool_helper(tool) assert mapped_tool is not None assert mapped_tool.get("defer_loading") is True @@ -1664,7 +1664,7 @@ def test_allowed_callers_field_preservation(): "allowed_callers": ["code_execution_20250825"], } - transformed_tool, _ = config._map_tool_helper(tool_with_allowed_callers) + transformed_tool, _ = config.map_tool_helper(tool_with_allowed_callers) assert transformed_tool is not None assert "allowed_callers" in transformed_tool assert transformed_tool["allowed_callers"] == ["code_execution_20250825"] @@ -1752,7 +1752,7 @@ def test_code_execution_20250825_tool_type(): tool = {"type": "code_execution_20250825", "name": "code_execution"} - transformed_tool, _ = config._map_tool_helper(tool) + transformed_tool, _ = config.map_tool_helper(tool) assert transformed_tool is not None assert transformed_tool["type"] == "code_execution_20250825" assert transformed_tool["name"] == "code_execution" @@ -1777,7 +1777,7 @@ def test_allowed_callers_in_function_field(): }, } - transformed_tool, _ = config._map_tool_helper(tool) + transformed_tool, _ = config.map_tool_helper(tool) assert transformed_tool is not None assert "allowed_callers" in transformed_tool assert transformed_tool["allowed_callers"] == ["code_execution_20250825"] @@ -1808,7 +1808,7 @@ def test_input_examples_field_preservation(): ], } - transformed_tool, _ = config._map_tool_helper(tool_with_examples) + transformed_tool, _ = config.map_tool_helper(tool_with_examples) assert transformed_tool is not None assert "input_examples" in transformed_tool assert len(transformed_tool["input_examples"]) == 2 @@ -1866,7 +1866,7 @@ def test_input_examples_in_function_field(): }, } - transformed_tool, _ = config._map_tool_helper(tool) + transformed_tool, _ = config.map_tool_helper(tool) assert transformed_tool is not None assert "input_examples" in transformed_tool assert len(transformed_tool["input_examples"]) == 2 @@ -1893,7 +1893,7 @@ def test_input_examples_with_other_features(): "allowed_callers": ["code_execution_20250825"], } - transformed_tool, _ = config._map_tool_helper(tool) + transformed_tool, _ = config.map_tool_helper(tool) assert transformed_tool is not None assert "input_examples" in transformed_tool assert "defer_loading" in transformed_tool @@ -1921,7 +1921,7 @@ def test_input_examples_empty_list_not_added(): "input_examples": [], } - transformed_tool, _ = config._map_tool_helper(tool) + transformed_tool, _ = config.map_tool_helper(tool) assert transformed_tool is not None # Empty list should not be added assert "input_examples" not in transformed_tool or len(transformed_tool.get("input_examples", [])) == 0 @@ -2195,7 +2195,7 @@ def test_anthropic_drop_params_false_forwards_to_unsupported_model(): ], ) def test_anthropic_model_supports_effort_param_recognizes_supporting_models(model): - assert AnthropicConfig._model_supports_effort_param(model, "anthropic") is True + assert AnthropicConfig.model_supports_effort_param(model, "anthropic") is True @pytest.mark.parametrize( @@ -2208,7 +2208,7 @@ def test_anthropic_model_supports_effort_param_recognizes_supporting_models(mode ], ) def test_anthropic_model_supports_effort_param_rejects_non_supporting_models(model): - assert AnthropicConfig._model_supports_effort_param(model, "anthropic") is False + assert AnthropicConfig.model_supports_effort_param(model, "anthropic") is False @pytest.mark.parametrize( @@ -2532,7 +2532,7 @@ def test_supports_effort_level_handles_provider_prefixes(model, level, expected) ], ) def test_validate_effort_for_model_centralises_per_model_gating(model, effort, expect_error): - err = AnthropicConfig._validate_effort_for_model(model, effort, "anthropic") + err = AnthropicConfig.validate_effort_for_model(model, effort, "anthropic") if expect_error: assert err is not None assert effort in err @@ -2807,7 +2807,7 @@ def test_is_adaptive_thinking_model_is_sourced_from_cost_map(local_model_cost_ma fallback for ids the cost map cannot resolve. The dated Claude 4.0 names stay non-adaptive because the date suffix is not read as a minor version, while 4.8/4.9/5.x are covered without a code change.""" - assert AnthropicConfig._is_adaptive_thinking_model(model, "anthropic") is expected + assert AnthropicConfig.is_adaptive_thinking_model(model, "anthropic") is expected def test_get_supported_params_includes_reasoning_for_sonnet_4_6_alias( @@ -4413,7 +4413,7 @@ def test_map_tool_helper_enforces_object_type_when_missing(): } original_params = tool["function"]["parameters"].copy() - result, _ = config._map_tool_helper(tool) + result, _ = config.map_tool_helper(tool) assert result is not None assert result["input_schema"]["type"] == "object" assert "properties" in result["input_schema"] @@ -4444,7 +4444,7 @@ def test_map_tool_helper_enforces_object_type_when_wrong_type(): } original_params = tool["function"]["parameters"].copy() - result, _ = config._map_tool_helper(tool) + result, _ = config.map_tool_helper(tool) assert result is not None assert result["input_schema"]["type"] == "object" assert result["input_schema"].get("properties") == {}, ( @@ -4478,7 +4478,7 @@ def test_map_tool_helper_preserves_valid_object_schema(): }, } - result, _ = config._map_tool_helper(tool) + result, _ = config.map_tool_helper(tool) assert result is not None assert result["input_schema"]["type"] == "object" assert "city" in result["input_schema"]["properties"] @@ -4500,7 +4500,7 @@ def test_map_tool_helper_empty_parameters_get_default(): }, } - result, _ = config._map_tool_helper(tool) + result, _ = config.map_tool_helper(tool) assert result is not None assert result["input_schema"]["type"] == "object" assert result["input_schema"].get("properties") == {} @@ -4559,7 +4559,7 @@ def test_advisor_tool_map_tool_helper(): "name": "advisor", "model": "claude-opus-4-6", } - returned_tool, mcp_server = config._map_tool_helper(tool) # type: ignore + returned_tool, mcp_server = config.map_tool_helper(tool) # type: ignore assert returned_tool is not None assert returned_tool["type"] == "advisor_20260301" assert returned_tool["model"] == "claude-opus-4-6" @@ -4576,7 +4576,7 @@ def test_advisor_tool_map_tool_helper_with_optional_fields(): "max_uses": 3, "caching": {"type": "ephemeral", "ttl": "5m"}, } - returned_tool, _ = config._map_tool_helper(tool) # type: ignore + returned_tool, _ = config.map_tool_helper(tool) # type: ignore assert returned_tool is not None assert returned_tool["max_uses"] == 3 assert returned_tool["caching"] == {"type": "ephemeral", "ttl": "5m"} @@ -4587,7 +4587,7 @@ def test_advisor_tool_map_tool_helper_missing_model(): config = AnthropicConfig() tool = {"type": "advisor_20260301", "name": "advisor"} with pytest.raises(ValueError, match="valid model"): - config._map_tool_helper(tool) # type: ignore + config.map_tool_helper(tool) # type: ignore def test_advisor_beta_header_injected(): @@ -5515,7 +5515,7 @@ def test_map_tool_helper_inlines_components_schemas_refs(): }, } - transformed, _ = config._map_tool_helper(tool) + transformed, _ = config.map_tool_helper(tool) assert transformed is not None schema = transformed["input_schema"] @@ -5555,7 +5555,7 @@ def test_map_tool_helper_inlines_legacy_definitions_refs(): }, } - transformed, _ = config._map_tool_helper(tool) + transformed, _ = config.map_tool_helper(tool) assert transformed is not None schema = transformed["input_schema"] @@ -5586,7 +5586,7 @@ def test_map_tool_helper_preserves_native_dollar_defs(): }, } - transformed, _ = config._map_tool_helper(tool) + transformed, _ = config.map_tool_helper(tool) assert transformed is not None schema = transformed["input_schema"] @@ -5618,7 +5618,7 @@ def test_map_tool_helper_does_not_mutate_caller_dict(): } snapshot = copy.deepcopy(tool) - config._map_tool_helper(tool) + config.map_tool_helper(tool) assert tool == snapshot, "caller's tool dict was mutated in place" @@ -5658,7 +5658,7 @@ def test_map_tool_helper_collision_prefers_definitions_over_components_schemas() }, } - transformed, _ = config._map_tool_helper(tool) + transformed, _ = config.map_tool_helper(tool) assert transformed is not None expected = {"type": "string", "description": "from-definitions"} @@ -5856,7 +5856,7 @@ def test_namespace_tool_flat_nested_tools_are_extracted(): ], } ] - anthropic_tools, _ = config._map_tools(tools) + anthropic_tools, _ = config.map_tools(tools) assert len(anthropic_tools) == 1 assert anthropic_tools[0]["name"] == "close_agent" @@ -5907,7 +5907,7 @@ def test_namespace_tool_nested_tools_are_extracted(): }, }, ] - anthropic_tools, mcp_servers = config._map_tools(tools) + anthropic_tools, mcp_servers = config.map_tools(tools) names = [t["name"] for t in anthropic_tools] assert "close_agent" in names assert "resume_agent" in names @@ -6323,7 +6323,7 @@ def _eager_chat_tool(**extra: object) -> dict[str, object]: @pytest.mark.parametrize("flag", [True, False]) def test_eager_input_streaming_passed_through_from_tool_top_level(flag): - mapped_tool, _ = AnthropicConfig()._map_tool_helper(_eager_chat_tool(eager_input_streaming=flag)) + mapped_tool, _ = AnthropicConfig().map_tool_helper(_eager_chat_tool(eager_input_streaming=flag)) assert mapped_tool == { "name": "write_file", @@ -6335,7 +6335,7 @@ def test_eager_input_streaming_passed_through_from_tool_top_level(flag): def test_eager_input_streaming_passed_through_from_function(): - mapped_tool, _ = AnthropicConfig()._map_tool_helper( + mapped_tool, _ = AnthropicConfig().map_tool_helper( {"type": "function", "function": _eager_chat_function(eager_input_streaming=True)} ) @@ -6344,14 +6344,14 @@ def test_eager_input_streaming_passed_through_from_function(): def test_eager_input_streaming_absent_stays_absent(): - mapped_tool, _ = AnthropicConfig()._map_tool_helper(_eager_chat_tool()) + mapped_tool, _ = AnthropicConfig().map_tool_helper(_eager_chat_tool()) assert "eager_input_streaming" not in mapped_tool def test_eager_input_streaming_rejects_non_boolean(): with pytest.raises(litellm.BadRequestError, match="eager_input_streaming must be a boolean"): - AnthropicConfig()._map_tool_helper(_eager_chat_tool(eager_input_streaming="true")) + AnthropicConfig().map_tool_helper(_eager_chat_tool(eager_input_streaming="true")) def test_eager_input_streaming_not_set_on_computer_use_tool(): @@ -6361,7 +6361,7 @@ def test_eager_input_streaming_not_set_on_computer_use_tool(): "eager_input_streaming": True, } - mapped_tool, _ = AnthropicConfig()._map_tool_helper(computer_tool) + mapped_tool, _ = AnthropicConfig().map_tool_helper(computer_tool) assert mapped_tool["type"] == "computer_20250124" assert "eager_input_streaming" not in mapped_tool diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/unit/llms/anthropic/pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index 06aaa4e61fb..2eee8549c0b 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -188,9 +188,7 @@ def test_translate_streaming_openai_chunk_to_anthropic_content_block(): ( block_type, content_block_start, - ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) print(content_block_start) @@ -236,9 +234,7 @@ def test_translate_streaming_openai_chunk_strips_gemini_thought_from_tool_call_i ( block_type, content_block_start, - ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "tool_use" assert content_block_start["id"] == base @@ -283,9 +279,7 @@ def test_translate_streaming_openai_chunk_to_anthropic_thinking_content_block(): ( block_type, content_block_start, - ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "thinking" assert content_block_start == { @@ -321,9 +315,7 @@ def test_translate_streaming_openai_chunk_to_anthropic_reasoning_content_only_co ( block_type, content_block_start, - ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "thinking" assert content_block_start == { @@ -369,9 +361,7 @@ def test_translate_streaming_openai_chunk_to_anthropic_thinking_signature_block( ( block_type, content_block_start, - ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "thinking" assert content_block_start == { @@ -424,9 +414,7 @@ def test_translate_streaming_openai_chunk_to_anthropic_content_block_thinking_an ( block_type, content_block_start, - ) = LiteLLMAnthropicMessagesAdapter()._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "thinking" @@ -1545,9 +1533,7 @@ def test_translate_streaming_openai_chunk_to_anthropic_emits_signature_when_thin ( block_type, content_block_start, - ) = adapter._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = adapter.translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "thinking" @@ -2040,9 +2026,7 @@ def test_streaming_chunk_with_both_text_and_tool_calls_issue_18238(): ( block_type, content_block_start, - ) = adapter._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = adapter.translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "tool_use" assert content_block_start["name"] == "Bash" @@ -2086,9 +2070,7 @@ def test_streaming_chunk_with_text_and_empty_tool_calls_returns_text_delta(): ( block_type, content_block_start, - ) = adapter._translate_streaming_openai_chunk_to_anthropic_content_block( - choices=choices - ) + ) = adapter.translate_streaming_openai_chunk_to_anthropic_content_block(choices=choices) assert block_type == "text" assert content_block_start == {"type": "text", "text": ""} @@ -3315,9 +3297,7 @@ def test_translate_openai_usage_to_anthropic_cache_tokens_from_dict_details_with "cache_write_tokens": 20.0, } - anthropic_usage = LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage_delta( - usage - ) + anthropic_usage = LiteLLMAnthropicMessagesAdapter.translate_openai_usage_to_anthropic_usage_delta(usage) assert anthropic_usage["input_tokens"] == 70 assert anthropic_usage["output_tokens"] == 50 @@ -3336,9 +3316,7 @@ def test_translate_openai_usage_to_anthropic_ignores_fractional_cache_tokens(): "cache_creation_tokens": 20.25, } - anthropic_usage = LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage_delta( - usage - ) + anthropic_usage = LiteLLMAnthropicMessagesAdapter.translate_openai_usage_to_anthropic_usage_delta(usage) assert anthropic_usage["input_tokens"] == 120 assert anthropic_usage["output_tokens"] == 50 @@ -3355,9 +3333,7 @@ def test_translate_openai_usage_to_anthropic_ignores_bool_cache_tokens(): usage.cache_read_input_tokens = True usage.cache_creation_input_tokens = True - anthropic_usage = LiteLLMAnthropicMessagesAdapter._translate_openai_usage_to_anthropic_usage_delta( - usage - ) + anthropic_usage = LiteLLMAnthropicMessagesAdapter.translate_openai_usage_to_anthropic_usage_delta(usage) assert anthropic_usage["input_tokens"] == 120 assert anthropic_usage["output_tokens"] == 50 diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py index e6963c33a56..1912876a453 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py @@ -15,11 +15,11 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( AnthropicMessagesStreamingResponse, BaseAnthropicMessagesStreamingIterator, _incomplete_stream_error_sse_event, - _is_message_stop_chunk, - _is_provider_error_chunk, anthropic_messages_response_as_sse_events, is_anthropic_content_delta_chunk, is_anthropic_ping_chunk, + is_message_stop_chunk, + is_provider_error_chunk, parse_anthropic_error_event, ) @@ -165,11 +165,11 @@ async def test_async_sse_wrapper_treats_message_stop_bytes_as_complete(): def test_is_message_stop_chunk(): - assert _is_message_stop_chunk({"type": "message_stop"}) is True - assert _is_message_stop_chunk({"type": "message_delta"}) is False - assert _is_message_stop_chunk(b"event: message_stop\ndata: {}\n\n") is True - assert _is_message_stop_chunk(b"raw-bytes") is False - assert _is_message_stop_chunk("message_stop") is False + assert is_message_stop_chunk({"type": "message_stop"}) is True + assert is_message_stop_chunk({"type": "message_delta"}) is False + assert is_message_stop_chunk(b"event: message_stop\ndata: {}\n\n") is True + assert is_message_stop_chunk(b"raw-bytes") is False + assert is_message_stop_chunk("message_stop") is False @pytest.mark.parametrize( @@ -202,7 +202,7 @@ def test_is_message_stop_chunk_ignores_substring_in_payload(): b'data: {"type": "content_block_delta", "delta": ' b'{"type": "input_json_delta", "partial_json": "\\"message_stop\\""}}\n\n' ) - assert _is_message_stop_chunk(delta_frame_with_substring) is False + assert is_message_stop_chunk(delta_frame_with_substring) is False def test_parse_anthropic_error_event_from_dict_chunk(): @@ -210,7 +210,7 @@ def test_parse_anthropic_error_event_from_dict_chunk(): (type, message, status) so the Router can decide whether to fall back.""" chunk = {"type": "error", "error": {"type": "overloaded_error", "message": "Overloaded"}} assert parse_anthropic_error_event(chunk) == ("overloaded_error", "Overloaded", 503) - assert _is_provider_error_chunk(chunk) is True + assert is_provider_error_chunk(chunk) is True def test_parse_anthropic_error_event_from_sse_bytes(): @@ -218,11 +218,10 @@ def test_parse_anthropic_error_event_from_sse_bytes(): Anthropic/Bedrock passthrough forwards verbatim today) must parse identically to the dict shape so the Router can raise a fallback.""" sse_chunk = ( - b"event: error\n" - b'data: {"type": "error", "error": {"type": "internal_server_error", "message": "boom"}}\n\n' + b'event: error\ndata: {"type": "error", "error": {"type": "internal_server_error", "message": "boom"}}\n\n' ) assert parse_anthropic_error_event(sse_chunk) == ("internal_server_error", "boom", 500) - assert _is_provider_error_chunk(sse_chunk) is True + assert is_provider_error_chunk(sse_chunk) is True def test_parse_anthropic_error_event_defaults_status_for_unknown_type(): @@ -248,7 +247,7 @@ def test_decoded_sse_data_line_swallows_invalid_json(): must not be treated as an error event or raise, just be ignored.""" malformed_frame = b"event: error\ndata: {not valid json\n\n" assert parse_anthropic_error_event(malformed_frame) is None - assert _is_provider_error_chunk(malformed_frame) is False + assert is_provider_error_chunk(malformed_frame) is False class TestIsAnthropicContentDeltaChunk: @@ -281,7 +280,7 @@ class TestIsAnthropicContentDeltaChunk: ) def test_parse_anthropic_error_event_non_error_chunks_return_none(chunk): assert parse_anthropic_error_event(chunk) is None - assert _is_provider_error_chunk(chunk) is False + assert is_provider_error_chunk(chunk) is False def test_parse_anthropic_error_event_ignores_substring_in_payload(): diff --git a/tests/unit/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py index 054cd8fd742..089538882eb 100644 --- a/tests/unit/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/unit/llms/anthropic/test_anthropic_common_utils.py @@ -1994,7 +1994,7 @@ class TestClaudeOpus48AdaptiveThinking: def test_adaptive_thinking_detected_for_opus_4_8(self, local_model_cost_map, model): from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is True @pytest.mark.parametrize( "model", @@ -2009,7 +2009,7 @@ class TestClaudeOpus48AdaptiveThinking: def test_adaptive_thinking_detected_for_fable_5(self, local_model_cost_map, model): from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is True @pytest.mark.parametrize( "model", @@ -2042,7 +2042,7 @@ class TestClaudeOpus48AdaptiveThinking: version (``4.6`` -> ``4-6``).""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is True @pytest.mark.parametrize( "model", @@ -2061,7 +2061,7 @@ class TestClaudeOpus48AdaptiveThinking: from litellm.llms.anthropic.common_utils import AnthropicModelInfo assert model not in litellm.model_cost - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is False + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is False @pytest.mark.parametrize( "model", @@ -2087,7 +2087,7 @@ class TestClaudeOpus48AdaptiveThinking: from litellm.llms.anthropic.common_utils import AnthropicModelInfo assert model not in litellm.model_cost - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is True @pytest.mark.parametrize( "model", @@ -2108,7 +2108,7 @@ class TestClaudeOpus48AdaptiveThinking: from litellm.llms.anthropic.common_utils import AnthropicModelInfo assert model not in litellm.model_cost - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is False + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is False @pytest.mark.parametrize( "model", @@ -2117,7 +2117,7 @@ class TestClaudeOpus48AdaptiveThinking: def test_non_adaptive_models_not_detected(self, local_model_cost_map, model): from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is False + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is False class TestDefaultSuffixAdaptiveThinking: @@ -2140,7 +2140,7 @@ class TestDefaultSuffixAdaptiveThinking: def test_default_suffix_models_are_adaptive_thinking(self, local_model_cost_map, model: str) -> None: from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True, ( + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is True, ( f"{model} not classified as adaptive thinking. Check _model_map_lookup_candidates strips @default suffix." ) @@ -2172,12 +2172,12 @@ class TestCapabilityProbeUsesCallerProvider: import litellm from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock") is True + assert AnthropicModelInfo.is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock") is True monkeypatch.setitem(litellm.model_cost[self.BEDROCK_MODEL], "supports_adaptive_thinking", False) litellm.get_model_info.cache_clear() - assert AnthropicModelInfo._is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock") is False + assert AnthropicModelInfo.is_adaptive_thinking_model(self.BEDROCK_MODEL, "bedrock") is False def test_create_anthropic_model_list_response_shape(): diff --git a/tests/unit/llms/anthropic/test_anthropic_reasoning_effort.py b/tests/unit/llms/anthropic/test_anthropic_reasoning_effort.py index 288817dff07..573f586de1e 100644 --- a/tests/unit/llms/anthropic/test_anthropic_reasoning_effort.py +++ b/tests/unit/llms/anthropic/test_anthropic_reasoning_effort.py @@ -13,26 +13,26 @@ from litellm.llms.anthropic.chat.transformation import AnthropicConfig class TestMapReasoningEffort: def test_none_returns_none_for_opus_4_6(self): """reasoning_effort=None should return None for Opus 4.6, not adaptive.""" - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort=None, model="claude-opus-4-6", custom_llm_provider="anthropic" ) assert result is None def test_none_returns_none_for_other_models(self): """reasoning_effort=None should return None for non-Opus models.""" - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort=None, model="claude-4-sonnet-20250514", custom_llm_provider="anthropic" ) assert result is None def test_opus_4_6_returns_adaptive_for_low(self): - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort="low", model="claude-opus-4-6", custom_llm_provider="anthropic" ) assert result["type"] == "adaptive" def test_opus_4_6_returns_adaptive_for_high(self): - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort="high", model="claude-opus-4-6", custom_llm_provider="anthropic" ) assert result["type"] == "adaptive" @@ -42,20 +42,20 @@ class TestMapReasoningEffort: """Regression LIT-5714: adaptive thinking without ``display`` makes Anthropic return a blank thinking block, so reasoning_effort callers always got ``reasoning_content: ""``.""" - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort=effort, model="claude-opus-4-6", custom_llm_provider="anthropic" ) assert result["display"] == "summarized" def test_other_model_low_returns_enabled_with_budget(self): - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort="low", model="claude-4-sonnet-20250514", custom_llm_provider="anthropic" ) assert result["type"] == "enabled" assert "budget_tokens" in result def test_other_model_high_returns_enabled_with_budget(self): - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort="high", model="claude-4-sonnet-20250514", custom_llm_provider="anthropic" ) assert result["type"] == "enabled" @@ -63,14 +63,14 @@ class TestMapReasoningEffort: def test_none_string_returns_none_for_opus_4_6(self): """reasoning_effort='none' should return None for Opus 4.6.""" - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort="none", model="claude-opus-4-6", custom_llm_provider="anthropic" ) assert result is None def test_none_string_returns_none_for_other_models(self): """reasoning_effort='none' should return None for non-Opus models.""" - result = AnthropicConfig._map_reasoning_effort( + result = AnthropicConfig.map_reasoning_effort( reasoning_effort="none", model="claude-4-sonnet-20250514", custom_llm_provider="anthropic" ) assert result is None diff --git a/tests/unit/llms/azure/realtime/test_azure_realtime_handler.py b/tests/unit/llms/azure/realtime/test_azure_realtime_handler.py index 73f43ec8d8a..7a5ce8ab853 100644 --- a/tests/unit/llms/azure/realtime/test_azure_realtime_handler.py +++ b/tests/unit/llms/azure/realtime/test_azure_realtime_handler.py @@ -85,7 +85,7 @@ async def test_construct_url_default_beta_protocol(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-4o-realtime-preview", api_version="2024-10-01-preview", @@ -106,7 +106,7 @@ async def test_construct_url_beta_protocol_explicit(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-4o-realtime-preview", api_version="2024-10-01-preview", @@ -126,7 +126,7 @@ async def test_construct_url_ga_protocol(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-4o-realtime-preview", api_version="2024-10-01-preview", @@ -153,7 +153,7 @@ async def test_construct_url_forwards_transcription_intent_ga(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-realtime-whisper", api_version="2025-04-01-preview", @@ -176,7 +176,7 @@ async def test_construct_url_forwards_transcription_intent_ga_without_model_quer from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-realtime-whisper", api_version="2025-04-01-preview", @@ -195,7 +195,7 @@ async def test_construct_url_forwards_transcription_intent_beta(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="whisper-deploy", api_version="2024-10-01-preview", @@ -213,7 +213,7 @@ async def test_construct_url_encodes_intent_value(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-realtime-whisper", api_version="2025-04-01-preview", @@ -230,7 +230,7 @@ async def test_construct_url_no_intent_when_absent(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-4o-realtime-preview", api_version="2024-10-01-preview", @@ -248,7 +248,7 @@ async def test_construct_url_v1_protocol(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-4o-realtime-preview", api_version="2024-10-01-preview", @@ -268,7 +268,7 @@ async def test_construct_url_case_insensitive_protocol(protocol): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime handler = AzureOpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base="https://my-endpoint.openai.azure.com", model="gpt-realtime-deployment", api_version=None, diff --git a/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py index 47609261a25..0b0a26aa2c6 100644 --- a/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/unit/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -23,7 +23,7 @@ async def test_get_openai_compatible_provider_info(): api_base, dynamic_api_key, custom_llm_provider, - ) = config._get_openai_compatible_provider_info( + ) = config.get_openai_compatible_provider_info( model="azure_ai/gpt-4o-mini", api_base="https://my-base", api_key="my-key", @@ -57,7 +57,7 @@ def test_foundry_base_keeps_azure_ai_provider(model: str, api_base: str, expecte _, _, custom_llm_provider, - ) = config._get_openai_compatible_provider_info( + ) = config.get_openai_compatible_provider_info( model=model, api_base=api_base, api_key="my-key", diff --git a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py index bfcdd8da3d8..052c4c422f2 100644 --- a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py +++ b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_handler.py @@ -161,12 +161,10 @@ class TestAzureAnthropicChatCompletion: mock_make_sync_call.assert_called_once() assert result is not None - @patch("litellm.llms.custom_httpx.http_handler._get_httpx_client") + @patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") @patch("litellm.utils.ProviderConfigManager") @patch("litellm.llms.azure_ai.anthropic.handler.AzureAnthropicConfig") - def test_completion_non_streaming( - self, mock_azure_config, mock_provider_manager, mock_get_client - ): + def test_completion_non_streaming(self, mock_azure_config, mock_provider_manager, mock_get_client): # Note: decorators are applied in reverse order """Test completion without streaming""" handler = AzureAnthropicChatCompletion() diff --git a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py index b78b2d0d842..866bdf173a9 100644 --- a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py +++ b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py @@ -40,7 +40,7 @@ class TestAzureAnthropicMessagesConfig: api_key = "test-api-key" with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-api-key"} result, api_base = config.validate_anthropic_messages_environment( @@ -73,7 +73,7 @@ class TestAzureAnthropicMessagesConfig: litellm_params = {"api_key": "test-api-key"} with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-api-key"} result, api_base = config.validate_anthropic_messages_environment( @@ -99,7 +99,7 @@ class TestAzureAnthropicMessagesConfig: litellm_params = {"api_key": "test-api-key"} with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-api-key"} result, api_base = config.validate_anthropic_messages_environment( diff --git a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_transformation.py b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_transformation.py index ddbc168589a..1179b0dd5da 100644 --- a/tests/unit/llms/azure_ai/claude/test_azure_anthropic_transformation.py +++ b/tests/unit/llms/azure_ai/claude/test_azure_anthropic_transformation.py @@ -30,7 +30,7 @@ class TestAzureAnthropicConfig: api_key = "test-api-key" with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-api-key"} result = config.validate_environment( @@ -59,7 +59,7 @@ class TestAzureAnthropicConfig: api_key = "test-api-key" with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-api-key"} result = config.validate_environment( @@ -87,7 +87,7 @@ class TestAzureAnthropicConfig: api_key = "provided-api-key" with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "provided-api-key"} config.validate_environment( @@ -113,7 +113,7 @@ class TestAzureAnthropicConfig: litellm_params = {"api_key": "test-api-key"} with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-api-key"} with patch.object(config, "get_anthropic_headers", return_value={}): @@ -139,7 +139,7 @@ class TestAzureAnthropicConfig: litellm_params = {"api_key": "test-api-key"} with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-api-key"} with patch.object(config, "get_anthropic_headers", return_value={}): @@ -163,7 +163,7 @@ class TestAzureAnthropicConfig: litellm_params = {"api_key": "test-api-key"} with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = { "api-key": "test-api-key", @@ -255,7 +255,7 @@ class TestAzureAnthropicConfig: headers = {"api-key": "test-key"} with patch( - "litellm.llms.azure.common_utils.BaseAzureLLM._base_validate_azure_environment" + "litellm.llms.azure.common_utils.BaseAzureLLM.base_validate_azure_environment" ) as mock_validate: mock_validate.return_value = {"api-key": "test-key"} result = config.transform_request( @@ -484,4 +484,3 @@ def test_chat_flagged_model_keeps_mid_conversation_system_role_in_place(local_mo "role": "system", "content": [{"type": "text", "text": "Answer with exactly one word."}], } - diff --git a/tests/unit/llms/bedrock/batches/test_handler.py b/tests/unit/llms/bedrock/batches/test_handler.py index e69098a460d..1cc8b58c768 100644 --- a/tests/unit/llms/bedrock/batches/test_handler.py +++ b/tests/unit/llms/bedrock/batches/test_handler.py @@ -1,4 +1,4 @@ -"""Unit tests for ``BedrockBatchesHandler._handle_model_invocation_job_status``. +"""Unit tests for ``BedrockBatchesHandler.handle_model_invocation_job_status``. These cover the upstream support for retrieving Bedrock bulk batch jobs (``arn:aws:bedrock:::model-invocation-job/``) — the ARN @@ -138,7 +138,7 @@ def test_predict_output_file_uri_returns_none_when_missing_input(missing_arg): def test_handle_model_invocation_job_status_completed(patched_boto3): fake_client, boto_client_factory = patched_boto3 - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) fake_client.get_model_invocation_job.assert_called_once_with(jobIdentifier=JOB_ARN) @@ -170,7 +170,7 @@ def test_completed_job_maps_provider_record_counts(patched_boto3, success_count, "errorRecordCount": error_count, } - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.request_counts is not None assert (batch.request_counts.total, batch.request_counts.completed, batch.request_counts.failed) == ( @@ -184,7 +184,7 @@ def test_missing_record_counts_leave_request_counts_none(patched_boto3): fake_client, _ = patched_boto3 fake_client.get_model_invocation_job.return_value = _fake_boto3_response() - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.request_counts is None @@ -193,7 +193,7 @@ def test_total_without_success_count_leaves_request_counts_none(patched_boto3): fake_client, _ = patched_boto3 fake_client.get_model_invocation_job.return_value = {**_fake_boto3_response(), "totalRecordCount": 100} - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.request_counts is None @@ -206,7 +206,7 @@ def test_missing_error_count_maps_to_zero_failed(patched_boto3): "successRecordCount": 100, } - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.request_counts is not None assert (batch.request_counts.total, batch.request_counts.completed, batch.request_counts.failed) == (100, 100, 0) @@ -232,11 +232,9 @@ def test_missing_error_count_maps_to_zero_failed(patched_boto3): ) def test_status_mapping(patched_boto3, bedrock_status, openai_status): fake_client, _ = patched_boto3 - fake_client.get_model_invocation_job.return_value = _fake_boto3_response( - status=bedrock_status - ) + fake_client.get_model_invocation_job.return_value = _fake_boto3_response(status=bedrock_status) - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.status == openai_status # output_file_id is only populated for terminal-completed jobs, so callers @@ -249,9 +247,7 @@ def test_status_mapping(patched_boto3, bedrock_status, openai_status): def test_explicit_region_overrides_arn(patched_boto3): _, boto_client_factory = patched_boto3 - BedrockBatchesHandler._handle_model_invocation_job_status( - batch_id=JOB_ARN, aws_region_name="eu-central-1" - ) + BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN, aws_region_name="eu-central-1") _, kwargs = boto_client_factory.call_args assert kwargs["region_name"] == "eu-central-1" @@ -262,7 +258,7 @@ def test_failure_message_propagates(patched_boto3): failed_response["message"] = "Input file failed validation" fake_client.get_model_invocation_job.return_value = failed_response - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.status == "failed" assert batch.failed_at == int(END_TIME.timestamp()) @@ -282,7 +278,7 @@ def test_completed_with_unpredictable_output_uri_stays_none(patched_boto3): incomplete_response["inputDataConfig"] = {"s3InputDataConfig": {"s3Uri": ""}} fake_client.get_model_invocation_job.return_value = incomplete_response - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.status == "completed" # output_file_id MUST be None (not the bare prefix) — that's the whole @@ -297,11 +293,9 @@ def test_completed_with_unpredictable_output_uri_stays_none(patched_boto3): def test_cancelled_status_sets_cancelled_at(patched_boto3): fake_client, _ = patched_boto3 - fake_client.get_model_invocation_job.return_value = _fake_boto3_response( - status="Stopped" - ) + fake_client.get_model_invocation_job.return_value = _fake_boto3_response(status="Stopped") - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.status == "cancelled" assert batch.cancelled_at == int(END_TIME.timestamp()) @@ -312,11 +306,9 @@ def test_cancelled_status_sets_cancelled_at(patched_boto3): def test_expired_status_sets_expired_at(patched_boto3): fake_client, _ = patched_boto3 - fake_client.get_model_invocation_job.return_value = _fake_boto3_response( - status="Expired" - ) + fake_client.get_model_invocation_job.return_value = _fake_boto3_response(status="Expired") - batch = BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + batch = BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) assert batch.status == "expired" assert batch.expired_at == int(END_TIME.timestamp()) @@ -331,18 +323,14 @@ def test_logging_obj_pre_and_post_call_invoked(patched_boto3): _, _ = patched_boto3 logging_obj = MagicMock() - BedrockBatchesHandler._handle_model_invocation_job_status( - batch_id=JOB_ARN, logging_obj=logging_obj - ) + BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN, logging_obj=logging_obj) logging_obj.pre_call.assert_called_once() logging_obj.post_call.assert_called_once() pre_kwargs = logging_obj.pre_call.call_args.kwargs assert pre_kwargs["input"] == JOB_ARN - assert pre_kwargs["additional_args"]["complete_input_dict"] == { - "jobIdentifier": JOB_ARN - } + assert pre_kwargs["additional_args"]["complete_input_dict"] == {"jobIdentifier": JOB_ARN} # Logged URL must use the bare job id, not the full ARN, so it doesn't # double the `model-invocation-job/` segment or embed colons in the path. assert pre_kwargs["additional_args"]["api_base"] == ( @@ -370,7 +358,7 @@ def test_missing_boto3_raises_helpful_import_error(): with patch("builtins.__import__", side_effect=fake_import): with pytest.raises(ImportError, match="pip install boto3"): - BedrockBatchesHandler._handle_model_invocation_job_status(batch_id=JOB_ARN) + BedrockBatchesHandler.handle_model_invocation_job_status(batch_id=JOB_ARN) def test_logging_url_uses_bare_id_when_only_id_passed(patched_boto3): @@ -379,7 +367,7 @@ def test_logging_url_uses_bare_id_when_only_id_passed(patched_boto3): _, _ = patched_boto3 logging_obj = MagicMock() - BedrockBatchesHandler._handle_model_invocation_job_status( + BedrockBatchesHandler.handle_model_invocation_job_status( batch_id=JOB_ID, aws_region_name="us-west-2", logging_obj=logging_obj ) @@ -530,7 +518,7 @@ def test_handle_model_invocation_job_status_builds_the_client_from_the_tagged_se return fake_bedrock with patch("boto3.client", side_effect=boto3_client): - batch = BedrockBatchesHandler._handle_model_invocation_job_status( + batch = BedrockBatchesHandler.handle_model_invocation_job_status( batch_id=JOB_ARN, aws_access_key_id="AKIABATCHSTATUSCALLER", aws_secret_access_key="pod-caller-secret", @@ -604,7 +592,7 @@ def test_retrieve_signs_with_deployment_credentials_when_env_bearer_token_is_set recorder: Final = _AuthorizationRecorder(_fake_boto3_response()) with patch("botocore.httpsession.URLLib3Session.send", recorder.send): - batch = BedrockBatchesHandler._handle_model_invocation_job_status( + batch = BedrockBatchesHandler.handle_model_invocation_job_status( batch_id=JOB_ARN, aws_access_key_id="AKIADEPLOYMENTKEY", aws_secret_access_key="deployment-secret", diff --git a/tests/unit/llms/bedrock/chat/test_converse_transformation.py b/tests/unit/llms/bedrock/chat/test_converse_transformation.py index 6db3836f68e..5b67266a586 100644 --- a/tests/unit/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/unit/llms/bedrock/chat/test_converse_transformation.py @@ -406,7 +406,7 @@ def test_transform_tool_call_with_cache_control(): }, ] - result = config.transform_request( + result = config._transform_request( model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"tools": tools}, @@ -2073,7 +2073,7 @@ async def test_transformation_directly(): messages = [{"role": "user", "content": "run ls command and find all python files"}] # Transform request - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"tools": tools}, @@ -2177,7 +2177,7 @@ def test_transform_request_with_multiple_tools(): messages = [{"role": "user", "content": "run ls command and find all python files"}] # Transform request - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"tools": tools}, @@ -2237,7 +2237,7 @@ def test_transform_request_with_computer_tool_only(): ] # Transform request - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"tools": tools}, @@ -2273,7 +2273,7 @@ def test_transform_request_with_bash_tool_only(): messages = [{"role": "user", "content": "run ls command and find all python files"}] # Transform request - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"tools": tools}, @@ -2309,7 +2309,7 @@ def test_transform_request_with_text_editor_tool(): messages = [{"role": "user", "content": "Edit this text file"}] # Transform request - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"tools": tools}, @@ -2361,7 +2361,7 @@ def test_transform_request_with_function_tool(): ] # Transform request - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"tools": tools}, @@ -3755,7 +3755,7 @@ def test_request_metadata_transformation(): ] # Transform request with requestMetadata - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": request_metadata}, @@ -3781,7 +3781,7 @@ def test_request_metadata_validation(): } # Should not raise exception - config.transform_request( + config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": valid_metadata}, @@ -3793,7 +3793,7 @@ def test_request_metadata_validation(): too_many_items = {f"key_{i}": f"value_{i}" for i in range(17)} with pytest.raises(Exception, match="maximum of 16 items") as exc_info: - config.transform_request( + config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": too_many_items}, @@ -3815,7 +3815,7 @@ def test_request_metadata_key_constraints(): invalid_metadata = {long_key: "value"} with pytest.raises(Exception, match=r"(?i)key length|256 characters"): - config.transform_request( + config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": invalid_metadata}, @@ -3827,7 +3827,7 @@ def test_request_metadata_key_constraints(): invalid_metadata = {"": "value"} with pytest.raises(Exception, match=r"(?i)key length|empty"): - config.transform_request( + config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": invalid_metadata}, @@ -3847,7 +3847,7 @@ def test_request_metadata_value_constraints(): invalid_metadata = {"key": long_value} with pytest.raises(Exception, match=r"(?i)value length|256 characters"): - config.transform_request( + config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": invalid_metadata}, @@ -3859,7 +3859,7 @@ def test_request_metadata_value_constraints(): valid_metadata = {"key": ""} # Should not raise exception - config.transform_request( + config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": valid_metadata}, @@ -3882,7 +3882,7 @@ def test_request_metadata_character_pattern(): } # Should not raise exception - config.transform_request( + config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": valid_metadata}, @@ -3917,7 +3917,7 @@ def test_request_metadata_with_other_params(): ] # Transform request with multiple parameters including request_metadata - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={ @@ -3947,7 +3947,7 @@ def test_request_metadata_empty(): messages = [{"role": "user", "content": "Hello!"}] # Empty dict should be allowed - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={"requestMetadata": {}}, @@ -3966,7 +3966,7 @@ def test_request_metadata_not_provided(): messages = [{"role": "user", "content": "Hello!"}] # No requestMetadata provided - request_data = config.transform_request( + request_data = config._transform_request( model="anthropic.claude-haiku-4-5-20251001-v1:0", messages=messages, optional_params={}, @@ -5005,7 +5005,7 @@ def test_parallel_tool_calls_newer_model_adds_disable_flag(): drop_params=False, ) - request_data = config.transform_request( + request_data = config._transform_request( model=model, messages=messages, optional_params=optional_params, @@ -5044,7 +5044,7 @@ def test_parallel_tool_calls_flag_decoupled_from_ttl_pricing(monkeypatch): drop_params=False, ) - request_data = config.transform_request( + request_data = config._transform_request( model=model, messages=messages, optional_params=optional_params, @@ -5073,7 +5073,7 @@ def test_parallel_tool_calls_older_model_drops_disable_flag(): drop_params=False, ) - request_data = config.transform_request( + request_data = config._transform_request( model=model, messages=messages, optional_params=optional_params, @@ -5102,7 +5102,7 @@ def test_parallel_tool_calls_emits_typed_auto_tool_choice(parallel_tool_calls, e drop_params=False, ) - request_data = config.transform_request( + request_data = config._transform_request( model=model, messages=messages, optional_params=optional_params, @@ -5136,7 +5136,7 @@ def test_parallel_tool_calls_with_explicit_tool_choice_omits_conflicting_type(to drop_params=False, ) - request_data = config.transform_request( + request_data = config._transform_request( model=model, messages=messages, optional_params=optional_params, @@ -5159,7 +5159,7 @@ def test_tool_choice_type_kept_when_no_tool_config_choice_conflicts(): drop_params=False, ) - request_data = config.transform_request( + request_data = config._transform_request( model=model, messages=[{"role": "user", "content": "What's the weather in SF and NYC?"}], optional_params=optional_params, @@ -5786,7 +5786,7 @@ def test_cache_points_emitted_only_for_models_that_support_prompt_caching(model, monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) - body = AmazonConverseConfig().transform_request( + body = AmazonConverseConfig()._transform_request( model=model, messages=[ {"role": "system", "content": [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]}, @@ -6640,7 +6640,7 @@ def test_converse_top_k_dropped_for_models_that_removed_it(): transform must strip it for models that removed sampling params (#30064).""" config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-fable-5", messages=[{"role": "user", "content": "hello"}], optional_params={"top_k": 40}, @@ -6656,7 +6656,7 @@ def test_converse_top_k_raises_without_drop_params(monkeypatch): config = AmazonConverseConfig() with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"): - config.transform_request( + config._transform_request( model="us.anthropic.claude-fable-5", messages=[{"role": "user", "content": "hello"}], optional_params={"top_k": 40}, @@ -6668,7 +6668,7 @@ def test_converse_top_k_raises_without_drop_params(monkeypatch): def test_converse_top_k_forwarded_on_models_that_accept_it(): config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-sonnet-4-6", messages=[{"role": "user", "content": "hello"}], optional_params={"top_k": 40}, @@ -6687,7 +6687,7 @@ def test_converse_top_k_zero_raises_without_drop_params(monkeypatch): config = AmazonConverseConfig() with pytest.raises(litellm.utils.UnsupportedParamsError, match="drop_params"): - config.transform_request( + config._transform_request( model="us.anthropic.claude-fable-5", messages=[{"role": "user", "content": "hello"}], optional_params={"top_k": 0}, @@ -6699,7 +6699,7 @@ def test_converse_top_k_zero_raises_without_drop_params(monkeypatch): def test_converse_top_k_zero_forwarded_on_models_that_accept_it(): config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-sonnet-4-6", messages=[{"role": "user", "content": "hello"}], optional_params={"top_k": 0}, @@ -6935,7 +6935,7 @@ def test_transform_request_no_tools_with_tool_history_succeeds_24158(monkeypatch monkeypatch.setattr(litellm, "modify_params", False) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=_orphaned_tool_history_messages(), optional_params={}, @@ -6956,7 +6956,7 @@ def test_transform_request_tool_unsupported_model_no_toolconfig_27138(monkeypatc monkeypatch.setattr(litellm, "modify_params", True) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="meta.llama3-2-3b-instruct-v1:0", messages=_orphaned_tool_history_messages(), optional_params={}, @@ -6975,7 +6975,7 @@ def test_transform_request_empty_tools_with_tool_history(monkeypatch, tools_valu monkeypatch.setattr(litellm, "modify_params", False) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=_orphaned_tool_history_messages(), optional_params={"tools": tools_value}, @@ -6992,7 +6992,7 @@ def test_transform_request_tool_result_only_history(monkeypatch): monkeypatch.setattr(litellm, "modify_params", False) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=[ {"role": "user", "content": "hi"}, @@ -7015,7 +7015,7 @@ def test_transform_request_neutralized_tool_output_is_guarded(monkeypatch): monkeypatch.setattr(litellm, "modify_params", False) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=[ {"role": "user", "content": "look it up"}, @@ -7055,7 +7055,7 @@ def test_transform_request_neutralized_tool_output_guarded_mid_history(monkeypat monkeypatch.setattr(litellm, "modify_params", False) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=[ {"role": "user", "content": "look it up"}, @@ -7117,7 +7117,7 @@ def test_transform_request_with_tools_still_builds_toolconfig(monkeypatch): monkeypatch.setattr(litellm, "modify_params", False) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=_orphaned_tool_history_messages(), optional_params={ @@ -7147,7 +7147,7 @@ def test_transform_request_flag_off_restores_raise(monkeypatch): config = AmazonConverseConfig() with pytest.raises(litellm.utils.UnsupportedParamsError, match="without `tools="): - config.transform_request( + config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=_orphaned_tool_history_messages(), optional_params={}, @@ -7163,7 +7163,7 @@ def test_transform_request_flag_off_with_modify_params_restores_dummy_tool(monke monkeypatch.setattr(litellm, "modify_params", True) config = AmazonConverseConfig() - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=_orphaned_tool_history_messages(), optional_params={}, @@ -7181,7 +7181,7 @@ def test_transform_request_flag_on_is_default(monkeypatch): config = AmazonConverseConfig() assert litellm.bedrock_neutralize_orphaned_tool_blocks is True - result = config.transform_request( + result = config._transform_request( model="us.anthropic.claude-opus-4-5-20251101-v1:0", messages=_orphaned_tool_history_messages(), optional_params={}, @@ -7360,7 +7360,7 @@ def test_legacy_thinking_translated_to_adaptive_on_adaptive_only_converse(model, model=model, drop_params=False, ) - request = config.transform_request( + request = config._transform_request( model=model, messages=[{"role": "user", "content": "hi"}], optional_params=optional_params, @@ -8067,7 +8067,7 @@ def test_flagged_model_replays_a_byte_identical_prefix_around_a_mid_conversation user turn in place; hoisting it into ``system`` would change the prefix every signed thinking block in the history is bound to.""" requests = [ - AmazonConverseConfig().transform_request( + AmazonConverseConfig()._transform_request( model="bedrock/us.anthropic.claude-fable-5-1", messages=copy.deepcopy(turn), optional_params={}, diff --git a/tests/unit/llms/bedrock/chat/test_invoke_handler.py b/tests/unit/llms/bedrock/chat/test_invoke_handler.py index dfe1c06edb5..09e40c95eb9 100644 --- a/tests/unit/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/unit/llms/bedrock/chat/test_invoke_handler.py @@ -88,7 +88,7 @@ def test_transform_tool_calls_index(): decoder = AWSEventStreamDecoder(model="test") parsed_chunks = [] for chunk in chunks: - parsed_chunk = decoder._chunk_parser(chunk) + parsed_chunk = decoder.chunk_parser(chunk) parsed_chunks.append(parsed_chunk) tool_call_chunks1 = parsed_chunks[8:12] tool_call_chunks2 = parsed_chunks[13:17] @@ -179,7 +179,7 @@ def test_transform_tool_calls_index_with_optional_arg_func(): decoder = AWSEventStreamDecoder(model="test") parsed_chunks = [] for chunk in chunks: - parsed_chunk = decoder._chunk_parser(chunk) + parsed_chunk = decoder.chunk_parser(chunk) parsed_chunks.append(parsed_chunk) tool_call_chunks = parsed_chunks[11:14] for tool_call_hunk in tool_call_chunks: @@ -343,7 +343,7 @@ def _converse_stream_wrapper(events, model=CONVERSE_MODEL): async def bedrock_stream(): decoder = AWSEventStreamDecoder(model=model) for event in events: - yield decoder._chunk_parser(chunk_data=event) + yield decoder.chunk_parser(chunk_data=event) return CustomStreamWrapper( completion_stream=bedrock_stream(), diff --git a/tests/unit/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py b/tests/unit/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py index fbcbd0aaea6..da5993fff6d 100644 --- a/tests/unit/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py +++ b/tests/unit/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py @@ -47,9 +47,7 @@ class TestBedrockAsyncInvokeEmbedding: ) config = TwelveLabsMarengoEmbeddingConfig() - response = config._transform_async_invoke_response( - async_invoke_response, "test-model" - ) + response = config.transform_async_invoke_response(async_invoke_response, "test-model") # Verify response structure assert isinstance(response, litellm.EmbeddingResponse) @@ -121,18 +119,14 @@ class TestBedrockAsyncInvokeEmbedding: elif input_type == "image": input_data = test_image_base64 elif input_type in ["video", "audio"]: - input_data = ( - "s3://test-bucket/test-file.mp4" - if input_type == "video" - else "s3://test-bucket/test-file.wav" - ) + input_data = "s3://test-bucket/test-file.mp4" if input_type == "video" else "s3://test-bucket/test-file.wav" inference_params = { "inputType": input_type, # This will be set by the parameter mapping "output_s3_uri": "s3://test-bucket/async-invoke-output/", } - transformed_request = config._transform_request( + transformed_request = config.transform_request( input=input_data, inference_params=inference_params, async_invoke_route=True, @@ -274,7 +268,7 @@ class TestBedrockAsyncInvokeEmbedding: mock_status.return_value = async_invoke_status_response # This would be called internally, but we can test the method directly - status_response = await bedrock_embedding._get_async_invoke_status( + status_response = await bedrock_embedding.get_async_invoke_status( invocation_arn="arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456", aws_region_name="us-east-1", ) @@ -290,10 +284,8 @@ class TestBedrockAsyncInvokeEmbedding: config = TwelveLabsMarengoEmbeddingConfig() - with pytest.raises( - ValueError, match="output_s3_uri cannot be empty for async invoke requests" - ): - config._transform_request( + with pytest.raises(ValueError, match="output_s3_uri cannot be empty for async invoke requests"): + config.transform_request( input=test_input, inference_params={"inputType": "text"}, async_invoke_route=True, @@ -309,10 +301,8 @@ class TestBedrockAsyncInvokeEmbedding: config = TwelveLabsMarengoEmbeddingConfig() - with pytest.raises( - ValueError, match="Input type 'video' requires async_invoke route" - ): - config._transform_request( + with pytest.raises(ValueError, match="Input type 'video' requires async_invoke route"): + config.transform_request( input="s3://test-bucket/test-video.mp4", inference_params={"inputType": "video"}, async_invoke_route=False, # Should fail for video without async route @@ -338,9 +328,7 @@ class TestBedrockAsyncInvokeEmbedding: for arn in test_cases: mock_response = {"invocationArn": arn} - response = config._transform_async_invoke_response( - mock_response, "test-model" - ) + response = config.transform_async_invoke_response(mock_response, "test-model") assert response._hidden_params._invocation_arn == arn @@ -351,9 +339,7 @@ class TestBedrockAsyncInvokeEmbedding: ) config = TwelveLabsMarengoEmbeddingConfig() - response = config._transform_async_invoke_response( - async_invoke_response, "test-model" - ) + response = config.transform_async_invoke_response(async_invoke_response, "test-model") # Test that hidden params can be accessed like a dictionary assert ( @@ -449,7 +435,7 @@ async def test_async_invoke_status_signs_off_the_event_loop(monkeypatch): return_value=httpx.Response(200, json=async_invoke_status_response) ) release = asyncio.create_task(probe.release_refresh_from_the_loop()) - status = await embedder._get_async_invoke_status( + status = await embedder.get_async_invoke_status( invocation_arn=async_invoke_status_response["invocationArn"], aws_region_name="us-east-1" ) await release diff --git a/tests/unit/llms/bedrock/embed/test_bedrock_embedding.py b/tests/unit/llms/bedrock/embed/test_bedrock_embedding.py index ad21cadaa4b..2aa2c22e298 100644 --- a/tests/unit/llms/bedrock/embed/test_bedrock_embedding.py +++ b/tests/unit/llms/bedrock/embed/test_bedrock_embedding.py @@ -848,7 +848,7 @@ def test_titan_multimodal_embedding_image_cost_tracking(): # Simulate batch_data with an image request (inputImage key set by _transform_request) batch_data = [{"inputImage": "/9j/4AAQSkZJRg=="}] - result = config._transform_response( + result = config.transform_response( response_list=response_list, model="amazon.titan-embed-image-v1", batch_data=batch_data, @@ -877,7 +877,7 @@ def test_titan_multimodal_embedding_text_no_image_count(): # Text-only request — no inputImage key batch_data = [{"inputText": "hello world"}] - result = config._transform_response( + result = config.transform_response( response_list=response_list, model="amazon.titan-embed-image-v1", batch_data=batch_data, @@ -904,7 +904,7 @@ def test_titan_multimodal_embedding_backward_compat_no_batch_data(): ] # Call without batch_data — should not break - result = config._transform_response( + result = config.transform_response( response_list=response_list, model="amazon.titan-embed-image-v1", ) @@ -1291,13 +1291,16 @@ def test_marengo_2_7_embedding_keeps_the_flat_payload(): def test_marengo_usage_counts_text_requests_and_images_across_a_batch(): duck = {"mediaType": "image", "base64String": "ZHVjaw=="} - response = TwelveLabsMarengoEmbeddingConfig()._transform_response( + response = TwelveLabsMarengoEmbeddingConfig().transform_response( response_list=[marengo_3_embedding_response, marengo_3_embedding_response, marengo_3_embedding_response], model="us.twelvelabs.marengo-embed-3-0-v1:0", batch_data=[ {"inputType": "text", "text": {"inputText": "a duck"}}, {"inputType": "image", "image": {"mediaSource": {"base64String": "ZHVjaw=="}}}, - {"inputType": "multi_input", "multi_input": {"mediaSources": [{"name": "a", **duck}, {"name": "b", **duck}]}}, + { + "inputType": "multi_input", + "multi_input": {"mediaSources": [{"name": "a", **duck}, {"name": "b", **duck}]}, + }, ], ) @@ -1308,7 +1311,7 @@ def test_marengo_usage_counts_text_requests_and_images_across_a_batch(): def test_marengo_usage_without_request_data_bills_nothing(): - response = TwelveLabsMarengoEmbeddingConfig()._transform_response( + response = TwelveLabsMarengoEmbeddingConfig().transform_response( response_list=[marengo_3_embedding_response], model="us.twelvelabs.marengo-embed-3-0-v1:0" ) @@ -1318,7 +1321,7 @@ def test_marengo_usage_without_request_data_bills_nothing(): def test_marengo_response_items_without_an_embedding_are_skipped(): - response = TwelveLabsMarengoEmbeddingConfig()._transform_response( + response = TwelveLabsMarengoEmbeddingConfig().transform_response( response_list=[{"data": [{"embeddingOption": "visual-text", "startSec": 0.0}, {"embedding": [0.1, 0.2, 0.3]}]}], model="us.twelvelabs.marengo-embed-3-0-v1:0", ) diff --git a/tests/unit/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py b/tests/unit/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py index f149953b6f1..094da323f82 100644 --- a/tests/unit/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py +++ b/tests/unit/llms/bedrock/embed/test_twelvelabs_marengo_3_transformation.py @@ -201,10 +201,10 @@ def test_invalid_marengo_3_params_are_rejected_before_the_request_is_sent(params def test_config_sends_the_nested_payload_for_marengo_3_and_the_flat_one_for_2_7(): - nested = TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_US)._transform_request( + nested = TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_US).transform_request( input="hello", inference_params={"input_type": "text"} ) - flat = TwelveLabsMarengoEmbeddingConfig(model=MARENGO_27_US)._transform_request( + flat = TwelveLabsMarengoEmbeddingConfig(model=MARENGO_27_US).transform_request( input="hello", inference_params={"input_type": "text"} ) assert nested == {"inputType": "text", "text": {"inputText": "hello"}} @@ -212,20 +212,20 @@ def test_config_sends_the_nested_payload_for_marengo_3_and_the_flat_one_for_2_7( def test_config_without_a_model_keeps_the_2_7_payload(): - request = TwelveLabsMarengoEmbeddingConfig()._transform_request(input="hello", inference_params={}) + request = TwelveLabsMarengoEmbeddingConfig().transform_request(input="hello", inference_params={}) assert request == {"inputType": "text", "inputText": "hello", "textTruncate": "end"} @pytest.mark.parametrize("input_type", ["video", "audio"]) def test_marengo_3_video_and_audio_still_require_the_async_route(input_type): with pytest.raises(ValueError, match=f"Input type '{input_type}' requires async_invoke route"): - TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_BASE)._transform_request( + TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_BASE).transform_request( input="s3://media/clip.mp4", inference_params={"input_type": input_type} ) def test_marengo_3_async_invoke_wraps_the_nested_payload_with_the_base_model_id(): - request = TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_BASE)._transform_request( + request = TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_BASE).transform_request( input="s3://media/clip.mp4", inference_params={ "input_type": "video", @@ -252,7 +252,7 @@ def test_marengo_3_async_invoke_wraps_the_nested_payload_with_the_base_model_id( def test_marengo_3_async_invoke_requires_an_output_s3_uri(): with pytest.raises(ValueError, match="output_s3_uri cannot be empty"): - TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_BASE)._transform_request( + TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_BASE).transform_request( input="hello", inference_params={"input_type": "text"}, async_invoke_route=True, @@ -355,8 +355,8 @@ def test_drop_params_comes_from_the_call_or_the_global(monkeypatch): def test_config_drops_marengo_2_7_only_params_only_when_asked(): config = TwelveLabsMarengoEmbeddingConfig(model=MARENGO_3_US) with pytest.raises(BedrockError, match=r"Marengo 2\.7 parameters textTruncate"): - config._transform_request("hello", {"textTruncate": "end"}) - assert config._transform_request("hello", {"textTruncate": "end"}, drop_params=True) == { + config.transform_request("hello", {"textTruncate": "end"}) + assert config.transform_request("hello", {"textTruncate": "end"}, drop_params=True) == { "inputType": "text", "text": {"inputText": "hello"}, } diff --git a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py index 12275df404f..4bfac61545d 100644 --- a/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py +++ b/tests/unit/llms/bedrock/files/test_bedrock_files_transformation.py @@ -46,48 +46,36 @@ class TestBedrockFilesTransformation: openai_jsonl_content.append(json.loads(line)) # Transform the content - bedrock_jsonl_content = ( - transformation._transform_openai_jsonl_content_to_bedrock_jsonl_content( - openai_jsonl_content=openai_jsonl_content - ) + bedrock_jsonl_content = transformation._transform_openai_jsonl_content_to_bedrock_jsonl_content( + openai_jsonl_content=openai_jsonl_content ) # Basic validation - assert len(bedrock_jsonl_content) == len( - openai_jsonl_content - ), "Should have same number of records" + assert len(bedrock_jsonl_content) == len(openai_jsonl_content), "Should have same number of records" # Check structure of transformed records for i, record in enumerate(bedrock_jsonl_content): - assert "recordId" in record, f"Record {i+1} should have recordId" - assert "modelInput" in record, f"Record {i+1} should have modelInput" + assert "recordId" in record, f"Record {i + 1} should have recordId" + assert "modelInput" in record, f"Record {i + 1} should have modelInput" # Check recordId matches custom_id from input expected_custom_id = openai_jsonl_content[i].get("custom_id") - assert ( - record["recordId"] == expected_custom_id - ), f"Record {i+1} recordId should match custom_id" + assert record["recordId"] == expected_custom_id, f"Record {i + 1} recordId should match custom_id" # Check modelInput has expected structure model_input = record["modelInput"] - assert isinstance( - model_input, dict - ), f"Record {i+1} modelInput should be a dictionary" + assert isinstance(model_input, dict), f"Record {i + 1} modelInput should be a dictionary" # For Anthropic models, should have anthropic_version and messages if "anthropic.claude" in openai_jsonl_content[i]["body"]["model"]: - assert ( - "anthropic_version" in model_input - ), f"Record {i+1} should have anthropic_version" - assert "messages" in model_input, f"Record {i+1} should have messages" - assert ( - "max_tokens" in model_input - ), f"Record {i+1} should have max_tokens" + assert "anthropic_version" in model_input, f"Record {i + 1} should have anthropic_version" + assert "messages" in model_input, f"Record {i + 1} should have messages" + assert "max_tokens" in model_input, f"Record {i + 1} should have max_tokens" def test_batch_keeps_an_internal_prefixed_key_out_of_the_bedrock_model_input(self): from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - result: Final = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result: Final = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "internal-key-1", @@ -127,18 +115,14 @@ class TestBedrockFilesTransformation: "url": "/v1/chat/completions", "body": { "model": "us.amazon.nova-pro-v1:0", - "messages": [ - {"role": "user", "content": "What is the capital of France?"} - ], + "messages": [{"role": "user", "content": "What is the capital of France?"}], "max_tokens": 50, "temperature": 0.7, }, } ] - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( - openai_jsonl_content - ) + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content) assert len(result) == 1 record = result[0] @@ -198,29 +182,24 @@ class TestBedrockFilesTransformation: } ] - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( - openai_jsonl_content - ) + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content) assert len(result) == 1 model_input = result[0]["modelInput"] - assert ( - "additionalModelRequestFields" not in model_input - or model_input["additionalModelRequestFields"] - ), "additionalModelRequestFields must be absent or non-empty — Nova rejects {}" - assert ( - "system" not in model_input or model_input["system"] - ), "system must be absent or non-empty — Nova rejects []" + assert "additionalModelRequestFields" not in model_input or model_input["additionalModelRequestFields"], ( + "additionalModelRequestFields must be absent or non-empty — Nova rejects {}" + ) + assert "system" not in model_input or model_input["system"], ( + "system must be absent or non-empty — Nova rejects []" + ) # Validate the exact shape AWS accepts assert model_input == { "messages": [ { "role": "user", - "content": [ - {"text": "What is 1 + 1? Answer with just the number."} - ], + "content": [{"text": "What is 1 + 1? Answer with just the number."}], } ], "inferenceConfig": {"maxTokens": 16}, @@ -272,9 +251,7 @@ class TestBedrockFilesTransformation: } ] - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( - openai_jsonl_content - ) + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content) assert len(result) == 1 model_input = result[0]["modelInput"] @@ -339,9 +316,7 @@ class TestBedrockFilesTransformation: } ] - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( - openai_jsonl_content - ) + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content) assert len(result) == 1 model_input = result[0]["modelInput"] @@ -691,9 +666,7 @@ class TestBedrockFilesTransformation: } ] - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( - openai_jsonl_content - ) + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content) assert len(result) == 1 model_input = result[0]["modelInput"] @@ -713,7 +686,7 @@ class TestBedrockFilesTransformation: "bedrock/anthropic.claude-haiku-4-5-20251001-v1:0", ) - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "req-1", @@ -747,7 +720,7 @@ class TestBedrockFilesTransformation: "bedrock/amazon.titan-embed-text-v2:0", ) - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "embedding-1", @@ -770,7 +743,7 @@ class TestBedrockFilesTransformation: def test_unmapped_alias_falls_back_to_target_model(self): from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "req-1", @@ -804,7 +777,7 @@ class TestBedrockFilesTransformation: def test_record_provider_wins_over_target_model(self): from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "openai-1", @@ -832,7 +805,7 @@ class TestBedrockFilesTransformation: def test_embedding_alias_falls_back_to_target_model(self): from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "embedding-1", @@ -922,9 +895,7 @@ class TestBedrockFilesEmbeddingTransformation: with open(os.path.join(here, "expected_bedrock_batch_embeddings.jsonl")) as f: expected = [json.loads(line) for line in f if line.strip()] - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( - openai_jsonl - ) + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl) assert result == expected @@ -933,7 +904,7 @@ class TestBedrockFilesEmbeddingTransformation: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig config = BedrockFilesConfig() - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -954,7 +925,7 @@ class TestBedrockFilesEmbeddingTransformation: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig config = BedrockFilesConfig() - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -980,7 +951,7 @@ class TestBedrockFilesEmbeddingTransformation: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig config = BedrockFilesConfig() - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -999,7 +970,7 @@ class TestBedrockFilesEmbeddingTransformation: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig config = BedrockFilesConfig() - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -1023,7 +994,7 @@ class TestBedrockFilesEmbeddingTransformation: config = BedrockFilesConfig() with pytest.raises(ValueError, match="one input per JSONL record"): - config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -1045,7 +1016,7 @@ class TestBedrockFilesEmbeddingTransformation: config = BedrockFilesConfig() with pytest.raises(ValueError, match="missing required `input`"): - config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -1061,7 +1032,7 @@ class TestBedrockFilesEmbeddingTransformation: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig config = BedrockFilesConfig() - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "chat-1", @@ -1106,7 +1077,7 @@ class TestBedrockFilesEmbeddingTransformation: "bedrock/amazon.nova-2-multimodal-embeddings-v1:0", ): with pytest.raises(NotImplementedError, match="titan-embed-text-v2"): - config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -1128,7 +1099,7 @@ class TestBedrockFilesEmbeddingTransformation: "us.amazon.titan-embed-text-v2:0", "bedrock/us.amazon.titan-embed-text-v2:0", ): - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -1149,10 +1120,8 @@ class TestBedrockFilesEmbeddingTransformation: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig config = BedrockFilesConfig() - with pytest.raises( - (NotImplementedError, ValueError), match=r"pre-tokenized|one input per" - ): - config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + with pytest.raises((NotImplementedError, ValueError), match=r"pre-tokenized|one input per"): + config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -1174,7 +1143,7 @@ class TestBedrockFilesEmbeddingTransformation: config = BedrockFilesConfig() with pytest.raises(NotImplementedError, match="pre-tokenized"): - config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "e1", @@ -1193,7 +1162,7 @@ class TestBedrockFilesEmbeddingTransformation: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig config = BedrockFilesConfig() - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "ambiguous-1", @@ -1489,7 +1458,7 @@ class TestBedrockFilesEmbeddingTransformation: # just need to make sure we DON'T silently produce an inputText # body and call it a chat completion. config = BedrockFilesConfig() - result = config._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = config.transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "explicit-chat-with-input", @@ -1585,7 +1554,7 @@ class TestBedrockBatchNonChatEndpointRecords: def _transform(self, record: dict) -> dict: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content([record]) + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content([record]) assert len(result) == 1 assert result[0]["recordId"] == record["custom_id"] return result[0]["modelInput"] @@ -1837,7 +1806,7 @@ class TestBedrockBatchNonChatEndpointRecords: def test_mixed_endpoints_in_one_file_keep_their_own_shapes(self): from litellm.llms.bedrock.files.transformation import BedrockFilesConfig - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [ { "custom_id": "chat", @@ -1944,7 +1913,7 @@ class TestBedrockBatchAnthropicRowParams: from litellm.llms.bedrock.files.transformation import BedrockFilesConfig record = {"custom_id": "row-1", "method": "POST", "url": url, "body": {"model": model, **body}} - result = BedrockFilesConfig()._transform_openai_jsonl_content_to_bedrock_jsonl_content( + result = BedrockFilesConfig().transform_openai_jsonl_content_to_bedrock_jsonl_content( [record], target_model=target_model ) assert len(result) == 1 diff --git a/tests/unit/llms/bedrock/image/test_amazon_stability3_transformation.py b/tests/unit/llms/bedrock/image/test_amazon_stability3_transformation.py index dbde8565e13..8d3bea5c2a3 100644 --- a/tests/unit/llms/bedrock/image/test_amazon_stability3_transformation.py +++ b/tests/unit/llms/bedrock/image/test_amazon_stability3_transformation.py @@ -12,4 +12,4 @@ from litellm.llms.bedrock.image_generation.amazon_stability3_transformation impo def test_stability_image_core_is_v3_model(): model = "stability.stable-image-core-v1:1" - assert AmazonStability3Config._is_stability_3_model(model) + assert AmazonStability3Config.is_stability_3_model(model) diff --git a/tests/unit/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py b/tests/unit/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py index ab42f9ec312..d4b0bc5f7e4 100644 --- a/tests/unit/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py +++ b/tests/unit/llms/bedrock/invoke_agent/test_bedrock_agent_transformation.py @@ -98,9 +98,7 @@ class TestAmazonInvokeAgentConfig: litellm_params = {} headers = {} - result = config.transform_request( - model, sample_messages, optional_params, litellm_params, headers - ) + result = config.transform_request(model, sample_messages, optional_params, litellm_params, headers) expected = { "inputText": "What is the weather like?", diff --git a/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index a269d556262..52821f58482 100644 --- a/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -241,7 +241,7 @@ def test_chunk_parser_usage_transformation(): }, } - parsed = decoder._chunk_parser(chunk.copy()) # use copy to avoid side-effects + parsed = decoder.chunk_parser(chunk.copy()) # use copy to avoid side-effects # The invocation metrics key should be removed and replaced by `usage` assert "amazon-bedrock-invocationMetrics" not in parsed @@ -274,7 +274,7 @@ def test_chunk_parser_preserves_cache_usage_fields_with_invocation_metrics(): }, } - parsed = decoder._chunk_parser(chunk.copy()) + parsed = decoder.chunk_parser(chunk.copy()) assert "amazon-bedrock-invocationMetrics" not in parsed assert parsed["usage"]["cache_read_input_tokens"] == 9821 @@ -298,7 +298,7 @@ def test_chunk_parser_maps_cache_token_counts_from_invocation_metrics(): }, } - parsed = decoder._chunk_parser(chunk.copy()) + parsed = decoder.chunk_parser(chunk.copy()) assert parsed["usage"]["input_tokens"] == 10174 assert parsed["usage"]["output_tokens"] == 500 @@ -324,7 +324,7 @@ def test_chunk_parser_keeps_existing_token_counts_over_invocation_metrics(): }, } - parsed = decoder._chunk_parser(chunk.copy()) + parsed = decoder.chunk_parser(chunk.copy()) assert parsed["usage"]["input_tokens"] == 7 assert parsed["usage"]["output_tokens"] == 11 @@ -383,7 +383,7 @@ async def test_bedrock_sse_wrapper_preserves_cache_usage_with_invocation_metrics async def _decoded_stream(): # type: ignore[return-type] for chunk in raw_chunks: - yield decoder._chunk_parser(copy.deepcopy(chunk)) + yield decoder.chunk_parser(copy.deepcopy(chunk)) collected: list[bytes] = [] async for chunk in cfg.bedrock_sse_wrapper( @@ -1676,7 +1676,7 @@ def test_bedrock_messages_stream_decoder_keeps_safeguard_results(): tool_verdicts = {"toolu_01": {"type": "evaluated", "outcome": "not_flagged"}} safeguard_results = [{"type": "dangerous_tool_use", "status": {"type": "available", "tool_uses": tool_verdicts}}] - message_start = decoder._chunk_parser( + message_start = decoder.chunk_parser( { "type": "message_start", "message": { @@ -1695,7 +1695,7 @@ def test_bedrock_messages_stream_decoder_keeps_safeguard_results(): assert isinstance(message_start, dict) assert message_start["message"]["safeguard_results"] == safeguard_results - message_delta = decoder._chunk_parser( + message_delta = decoder.chunk_parser( { "type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None, "safeguard_results": safeguard_results}, diff --git a/tests/unit/llms/bedrock/test_base_aws_llm.py b/tests/unit/llms/bedrock/test_base_aws_llm.py index 0d88cc7dda3..f1c95aaf7a2 100644 --- a/tests/unit/llms/bedrock/test_base_aws_llm.py +++ b/tests/unit/llms/bedrock/test_base_aws_llm.py @@ -450,7 +450,7 @@ def test_get_aws_region_name_boto3_fallback(): mock_boto3_session.return_value = mock_session optional_params = {} - result = base_aws_llm._get_aws_region_name(optional_params) + result = base_aws_llm.get_aws_region_name(optional_params) assert result == "us-east-1" mock_boto3_session.assert_called_once() @@ -465,7 +465,7 @@ def test_get_aws_region_name_boto3_fallback(): mock_boto3_session.return_value = mock_session optional_params = {} - result = base_aws_llm._get_aws_region_name(optional_params) + result = base_aws_llm.get_aws_region_name(optional_params) assert result == "us-west-2" mock_boto3_session.assert_called_once() @@ -478,7 +478,7 @@ def test_get_aws_region_name_boto3_fallback(): mock_boto3_session.side_effect = Exception("boto3 not available") optional_params = {} - result = base_aws_llm._get_aws_region_name(optional_params) + result = base_aws_llm.get_aws_region_name(optional_params) assert result == "us-west-2" mock_boto3_session.assert_called_once() @@ -486,7 +486,7 @@ def test_get_aws_region_name_boto3_fallback(): # Test case 4: aws_region_name is provided in optional_params (should not use boto3) with patch("boto3.Session") as mock_boto3_session: optional_params = {"aws_region_name": "eu-west-1"} - result = base_aws_llm._get_aws_region_name(optional_params) + result = base_aws_llm.get_aws_region_name(optional_params) assert result == "eu-west-1" mock_boto3_session.assert_not_called() @@ -503,7 +503,7 @@ def test_get_aws_region_name_boto3_fallback(): with patch("boto3.Session") as mock_boto3_session: optional_params = {} - result = base_aws_llm._get_aws_region_name(optional_params) + result = base_aws_llm.get_aws_region_name(optional_params) assert result == "ap-southeast-1" mock_boto3_session.assert_not_called() @@ -534,9 +534,7 @@ def test_get_aws_region_name_rejects_malformed_region(bad_region): base_aws_llm = BaseAWSLLM() with pytest.raises(ValueError, match="Invalid AWS region format"): - base_aws_llm._get_aws_region_name( - optional_params={"aws_region_name": bad_region} - ) + base_aws_llm.get_aws_region_name(optional_params={"aws_region_name": bad_region}) @pytest.mark.parametrize( @@ -553,9 +551,7 @@ def test_get_aws_region_name_rejects_malformed_region(bad_region): def test_get_aws_region_name_accepts_valid_regions(valid_region): """Real AWS region formats must continue to work after the format guard.""" base_aws_llm = BaseAWSLLM() - result = base_aws_llm._get_aws_region_name( - optional_params={"aws_region_name": valid_region} - ) + result = base_aws_llm.get_aws_region_name(optional_params={"aws_region_name": valid_region}) assert result == valid_region @@ -576,7 +572,7 @@ def test_get_aws_region_name_rejects_malformed_region_from_env(): mock_get_secret.side_effect = side_effect with pytest.raises(ValueError, match="Invalid AWS region format"): - base_aws_llm._get_aws_region_name(optional_params={}) + base_aws_llm.get_aws_region_name(optional_params={}) def test_get_aws_region_name_for_non_llm_api_calls_rejects_malformed_param(): @@ -2704,29 +2700,19 @@ def test_converse_handler_external_id_extraction(): mock_credentials.token = "test-session-token" return mock_credentials - with patch.object( - converse_llm, "get_credentials", side_effect=mock_get_credentials - ): - with patch.object( - converse_llm, "_get_aws_region_name", return_value="us-west-2" - ): + with patch.object(converse_llm, "get_credentials", side_effect=mock_get_credentials): + with patch.object(converse_llm, "_get_aws_region_name", return_value="us-west-2"): with patch.object( converse_llm, "get_runtime_endpoint", return_value=("https://test", "https://test"), ): with patch("litellm.AmazonConverseConfig") as mock_config: - mock_config.return_value._transform_request.return_value = { - "test": "data" - } - with patch.object( - converse_llm, "get_request_headers" - ) as mock_headers: + mock_config.return_value._transform_request.return_value = {"test": "data"} + with patch.object(converse_llm, "get_request_headers") as mock_headers: mock_headers.return_value = MagicMock() mock_headers.return_value.headers = {"Authorization": "test"} - with patch( - "litellm.llms.custom_httpx.http_handler._get_httpx_client" - ) as mock_client: + with patch("litellm.llms.custom_httpx.http_handler.get_httpx_client") as mock_client: mock_http_client = MagicMock() mock_response = MagicMock() mock_response.raise_for_status.return_value = None @@ -2734,9 +2720,7 @@ def test_converse_handler_external_id_extraction(): mock_client.return_value = mock_http_client # Mock the transform_response method - mock_config.return_value._transform_response.return_value = ( - MagicMock() - ) + mock_config.return_value._transform_response.return_value = MagicMock() # Call completion with aws_external_id in optional_params optional_params = { diff --git a/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py index f9b183e10ec..9c35e3253e5 100644 --- a/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py +++ b/tests/unit/llms/bedrock_mantle/test_bedrock_mantle_transformation.py @@ -59,7 +59,7 @@ class TestBedrockMantleConfig: monkeypatch.setenv("BEDROCK_MANTLE_REGION", "eu-west-1") monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info(None, None) + api_base, _ = cfg.get_openai_compatible_provider_info(None, None) assert api_base == "https://bedrock-mantle.eu-west-1.api.aws/v1" def test_default_api_base_uses_aws_region(self, monkeypatch): @@ -67,7 +67,7 @@ class TestBedrockMantleConfig: monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) monkeypatch.setenv("AWS_REGION", "ap-northeast-1") cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info(None, None) + api_base, _ = cfg.get_openai_compatible_provider_info(None, None) assert api_base == "https://bedrock-mantle.ap-northeast-1.api.aws/v1" def test_default_api_base_uses_aws_region_name_env(self, monkeypatch): @@ -76,7 +76,7 @@ class TestBedrockMantleConfig: monkeypatch.delenv("AWS_REGION", raising=False) monkeypatch.setenv("AWS_REGION_NAME", "ca-central-1") cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info(None, None) + api_base, _ = cfg.get_openai_compatible_provider_info(None, None) assert api_base == "https://bedrock-mantle.ca-central-1.api.aws/v1" def test_aws_region_name_param_overrides_env(self, monkeypatch): @@ -85,7 +85,7 @@ class TestBedrockMantleConfig: monkeypatch.setenv("BEDROCK_MANTLE_REGION", "us-west-2") monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info( + api_base, _ = cfg.get_openai_compatible_provider_info( None, None, litellm_params=GenericLiteLLMParams(aws_region_name="us-east-2") ) assert api_base == "https://bedrock-mantle.us-east-2.api.aws/v1" @@ -98,12 +98,10 @@ class TestBedrockMantleConfig: monkeypatch.delenv("AWS_REGION", raising=False) cfg = BedrockMantleChatConfig() with pytest.raises(ValueError, match="api\\.aws\\.attacker\\.example/'\\. Region names must contain only"): - cfg._get_openai_compatible_provider_info( + cfg.get_openai_compatible_provider_info( None, None, - litellm_params=GenericLiteLLMParams( - aws_region_name="us-east-1.api.aws.attacker.example/" - ), + litellm_params=GenericLiteLLMParams(aws_region_name="us-east-1.api.aws.attacker.example/"), ) def test_get_llm_provider_rejects_malicious_aws_region_name(self, monkeypatch): @@ -143,7 +141,7 @@ class TestBedrockMantleConfig: for var in ("BEDROCK_MANTLE_REGION", "BEDROCK_MANTLE_API_BASE", "AWS_REGION", "AWS_REGION_NAME"): monkeypatch.delenv(var, raising=False) cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info(None, None, model="us-gov-west-1/xai.grok-4.3") + api_base, _ = cfg.get_openai_compatible_provider_info(None, None, model="us-gov-west-1/xai.grok-4.3") assert api_base == "https://bedrock-mantle.us-gov-west-1.api.aws/openai/v1" def test_aws_region_name_param_beats_model_region_prefix(self, monkeypatch, local_cost_map): @@ -152,7 +150,7 @@ class TestBedrockMantleConfig: for var in ("BEDROCK_MANTLE_REGION", "BEDROCK_MANTLE_API_BASE", "AWS_REGION", "AWS_REGION_NAME"): monkeypatch.delenv(var, raising=False) cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info( + api_base, _ = cfg.get_openai_compatible_provider_info( None, None, litellm_params=GenericLiteLLMParams(aws_region_name="us-east-1"), @@ -165,13 +163,13 @@ class TestBedrockMantleConfig: monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) monkeypatch.delenv("AWS_REGION", raising=False) cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info(None, None) + api_base, _ = cfg.get_openai_compatible_provider_info(None, None) assert api_base == "https://bedrock-mantle.us-east-1.api.aws/v1" def test_custom_api_base_overrides_default(self, monkeypatch): custom_base = "https://bedrock-mantle.us-west-2.api.aws/v1" cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info(custom_base, None) + api_base, _ = cfg.get_openai_compatible_provider_info(custom_base, None) assert api_base == custom_base def test_chat_base_for_gpt_oss_uses_v1(self, monkeypatch): @@ -181,18 +179,14 @@ class TestBedrockMantleConfig: monkeypatch.setenv("BEDROCK_MANTLE_REGION", "us-east-2") monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info( - None, None, model="openai.gpt-oss-120b" - ) + api_base, _ = cfg.get_openai_compatible_provider_info(None, None, model="openai.gpt-oss-120b") assert api_base == "https://bedrock-mantle.us-east-2.api.aws/v1" @pytest.mark.parametrize( "model_id", ["google.gemma-4-31b", "google.gemma-4-26b-a4b", "google.gemma-4-e2b"], ) - def test_chat_base_for_gemma_4_uses_openai_v1( - self, monkeypatch, local_cost_map, model_id - ): + def test_chat_base_for_gemma_4_uses_openai_v1(self, monkeypatch, local_cost_map, model_id): # The chat-config bug the Gemma 4 cards exposed: gemma-4-* is served on the # /openai/v1 base, not the hardcoded /v1. Driven by the price-map # use_openai_responses_path flag (loaded by local_cost_map). Fails before @@ -200,41 +194,35 @@ class TestBedrockMantleConfig: monkeypatch.setenv("BEDROCK_MANTLE_REGION", "us-east-2") monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info( - None, None, model=model_id - ) + api_base, _ = cfg.get_openai_compatible_provider_info(None, None, model=model_id) assert api_base == "https://bedrock-mantle.us-east-2.api.aws/openai/v1" - def test_chat_base_explicit_api_base_wins_over_derived( - self, monkeypatch, local_cost_map - ): + def test_chat_base_explicit_api_base_wins_over_derived(self, monkeypatch, local_cost_map): # An explicit api_base must not be overridden by the data-driven default, # even for a model whose default differs (gemma-4 -> openai/v1). monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) custom_base = "https://bedrock-mantle.us-west-2.api.aws/v1" cfg = BedrockMantleChatConfig() - api_base, _ = cfg._get_openai_compatible_provider_info( - custom_base, None, model="google.gemma-4-31b" - ) + api_base, _ = cfg.get_openai_compatible_provider_info(custom_base, None, model="google.gemma-4-31b") assert api_base == custom_base def test_api_key_from_env(self, monkeypatch): monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "test-key-123") cfg = BedrockMantleChatConfig() - _, api_key = cfg._get_openai_compatible_provider_info(None, None) + _, api_key = cfg.get_openai_compatible_provider_info(None, None) assert api_key == "test-key-123" def test_api_key_param_overrides_env(self, monkeypatch): monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "env-key") cfg = BedrockMantleChatConfig() - _, api_key = cfg._get_openai_compatible_provider_info(None, "explicit-key") + _, api_key = cfg.get_openai_compatible_provider_info(None, "explicit-key") assert api_key == "explicit-key" def test_api_key_from_aws_bearer_token_bedrock_env(self, monkeypatch): monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "standard-bearer") cfg = BedrockMantleChatConfig() - _, api_key = cfg._get_openai_compatible_provider_info(None, None) + _, api_key = cfg.get_openai_compatible_provider_info(None, None) assert api_key == "standard-bearer" def test_get_supported_openai_params(self): diff --git a/tests/unit/llms/chat/test_converse_handler.py b/tests/unit/llms/chat/test_converse_handler.py index 57bb9ab771f..6279b74bc0d 100644 --- a/tests/unit/llms/chat/test_converse_handler.py +++ b/tests/unit/llms/chat/test_converse_handler.py @@ -10,7 +10,7 @@ import pytest import litellm from litellm.llms.bedrock.chat import BedrockConverseLLM from litellm.llms.bedrock.chat.converse_handler import make_sync_call -from litellm.llms.bedrock.common_utils import _get_all_bedrock_regions +from litellm.llms.bedrock.common_utils import get_all_bedrock_regions from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from tests._support.stream_chunk_size import DEFAULT_CHUNKING_REQUESTS, ROUTER_CHUNK_SIZE_CASES, keys_at_every_depth @@ -90,7 +90,7 @@ class TestBedrockRegionInModelPath: _region_from_model = None _potential_region = _stripped.split("/", 1)[0] - if _potential_region in _get_all_bedrock_regions() and "/" in _stripped: + if _potential_region in get_all_bedrock_regions() and "/" in _stripped: _region_from_model = _potential_region _stripped = _stripped.split("/", 1)[1] _model_for_id = _stripped @@ -127,7 +127,7 @@ class TestBedrockRegionInModelPath: _stripped = model _region_from_model = None _potential_region = _stripped.split("/", 1)[0] - if _potential_region in _get_all_bedrock_regions() and "/" in _stripped: + if _potential_region in get_all_bedrock_regions() and "/" in _stripped: _region_from_model = _potential_region _stripped = _stripped.split("/", 1)[1] _model_for_id = _stripped diff --git a/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py b/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py index d367e4df0ab..43a4c26b6d2 100644 --- a/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py +++ b/tests/unit/llms/chatgpt/test_chatgpt_authenticator.py @@ -81,7 +81,7 @@ class TestChatGPTAuthenticator: client.post.return_value = response with patch( # test-quality-ok: requested seam for asserting timeout propagation - "litellm.llms.chatgpt.authenticator._get_httpx_client", return_value=client + "litellm.llms.chatgpt.authenticator.get_httpx_client", return_value=client ): refreshed = authenticator._refresh_tokens("refresh-123") diff --git a/tests/unit/llms/cohere/embed/test_v1_transformation.py b/tests/unit/llms/cohere/embed/test_v1_transformation.py index 66129b64a2c..31b4450ed88 100644 --- a/tests/unit/llms/cohere/embed/test_v1_transformation.py +++ b/tests/unit/llms/cohere/embed/test_v1_transformation.py @@ -31,7 +31,7 @@ class TestCohereEmbeddingV1Transform: data = {"texts": input_data, "input_type": "search_query"} model_response = EmbeddingResponse() - result = self.config._transform_response( + result = self.config.transform_response( response=mock_response, api_key="test-api-key", logging_obj=self.logging_obj, @@ -92,7 +92,7 @@ class TestCohereEmbeddingV1Transform: } model_response = EmbeddingResponse() - result = self.config._transform_response( + result = self.config.transform_response( response=mock_response, api_key="test-api-key", logging_obj=self.logging_obj, @@ -152,7 +152,7 @@ class TestCohereEmbeddingV1Transform: data = {"images": input_data, "input_type": "image"} model_response = EmbeddingResponse() - result = self.config._transform_response( + result = self.config.transform_response( response=mock_response, api_key="test-api-key", logging_obj=self.logging_obj, @@ -188,7 +188,7 @@ class TestCohereEmbeddingV1Transform: data = {"texts": input_data, "input_type": "search_query"} model_response = EmbeddingResponse() - result = self.config._transform_response( + result = self.config.transform_response( response=mock_response, api_key="test-api-key", logging_obj=self.logging_obj, diff --git a/tests/unit/llms/crusoe/test_crusoe.py b/tests/unit/llms/crusoe/test_crusoe.py index 718d00222aa..1fd2496cf3c 100644 --- a/tests/unit/llms/crusoe/test_crusoe.py +++ b/tests/unit/llms/crusoe/test_crusoe.py @@ -24,7 +24,7 @@ def test_crusoe_dynamic_config_defaults(): config = create_config_class(JSONProviderRegistry.get("crusoe"))() with patch.dict(os.environ, {}, clear=True): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == CRUSOE_API_BASE assert api_key is None @@ -41,7 +41,7 @@ def test_crusoe_dynamic_config_env_vars(): os.environ, {"CRUSOE_API_KEY": "test-key", "CRUSOE_API_BASE": "https://custom.crusoe.com/v1"}, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://custom.crusoe.com/v1" assert api_key == "test-key" @@ -55,9 +55,7 @@ def test_crusoe_dynamic_config_explicit_params(): config = create_config_class(JSONProviderRegistry.get("crusoe"))() with patch.dict(os.environ, {"CRUSOE_API_KEY": "env-key"}): - api_base, api_key = config._get_openai_compatible_provider_info( - "https://override.crusoe.com/v1", "override-key" - ) + api_base, api_key = config.get_openai_compatible_provider_info("https://override.crusoe.com/v1", "override-key") assert api_base == "https://override.crusoe.com/v1" assert api_key == "override-key" diff --git a/tests/unit/llms/custom_httpx/test_aiohttp_cleanup_closed.py b/tests/unit/llms/custom_httpx/test_aiohttp_cleanup_closed.py index 82010e82cea..03f83903e55 100644 --- a/tests/unit/llms/custom_httpx/test_aiohttp_cleanup_closed.py +++ b/tests/unit/llms/custom_httpx/test_aiohttp_cleanup_closed.py @@ -8,16 +8,13 @@ def test_create_aiohttp_transport_sets_enable_cleanup_closed_when_needed(monkeyp session_mock = MagicMock(name="session") monkeypatch.setattr(http_handler_module, "AIOHTTP_NEEDS_CLEANUP_CLOSED", True) - with patch.object( - http_handler_module, "TCPConnector", return_value=connector_mock - ) as mock_tcp_connector: - with patch.object( - http_handler_module, "ClientSession", return_value=session_mock - ), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")): - transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport( - shared_session=None - ) - transport._get_valid_client_session() + with patch.object(http_handler_module, "TCPConnector", return_value=connector_mock) as mock_tcp_connector: + with ( + patch.object(http_handler_module, "ClientSession", return_value=session_mock), + patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")), + ): + transport = http_handler_module.AsyncHTTPHandler.create_aiohttp_transport(shared_session=None) + transport.get_valid_client_session() assert mock_tcp_connector.call_args.kwargs["enable_cleanup_closed"] is True @@ -31,15 +28,12 @@ def test_create_aiohttp_transport_omits_enable_cleanup_closed_when_not_needed( session_mock = MagicMock(name="session") monkeypatch.setattr(http_handler_module, "AIOHTTP_NEEDS_CLEANUP_CLOSED", False) - with patch.object( - http_handler_module, "TCPConnector", return_value=connector_mock - ) as mock_tcp_connector: - with patch.object( - http_handler_module, "ClientSession", return_value=session_mock - ), patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")): - transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport( - shared_session=None - ) - transport._get_valid_client_session() + with patch.object(http_handler_module, "TCPConnector", return_value=connector_mock) as mock_tcp_connector: + with ( + patch.object(http_handler_module, "ClientSession", return_value=session_mock), + patch.object(http_handler_module, "DummyCookieJar", return_value=MagicMock(name="cookie_jar")), + ): + transport = http_handler_module.AsyncHTTPHandler.create_aiohttp_transport(shared_session=None) + transport.get_valid_client_session() assert "enable_cleanup_closed" not in mock_tcp_connector.call_args.kwargs diff --git a/tests/unit/llms/custom_httpx/test_aiohttp_handler.py b/tests/unit/llms/custom_httpx/test_aiohttp_handler.py index c58e6d6cf5c..83c38b65beb 100644 --- a/tests/unit/llms/custom_httpx/test_aiohttp_handler.py +++ b/tests/unit/llms/custom_httpx/test_aiohttp_handler.py @@ -288,16 +288,14 @@ class TestBaseLLMAIOHTTPHandler: mock_session_from_transport = Mock() mock_transport = Mock(spec=LiteLLMAiohttpTransport) - mock_transport._get_valid_client_session = Mock( - return_value=mock_session_from_transport - ) + mock_transport.get_valid_client_session = Mock(return_value=mock_session_from_transport) handler = BaseLLMAIOHTTPHandler(transport=mock_transport) result = handler._create_client_session_with_transport() # Should use transport's session creation method - mock_transport._get_valid_client_session.assert_called_once() + mock_transport.get_valid_client_session.assert_called_once() assert result is mock_session_from_transport # Should not call aiohttp.ClientSession directly @@ -433,21 +431,17 @@ class TestBaseLLMAIOHTTPHandler: def test_transport_priority_hierarchy(self): """Test that session creation follows the right priority: transport > connector > default""" - # Test with transport having _get_valid_client_session + # Test with transport having get_valid_client_session mock_transport = Mock(spec=LiteLLMAiohttpTransport) mock_session_from_transport = Mock() - mock_transport._get_valid_client_session = Mock( - return_value=mock_session_from_transport - ) + mock_transport.get_valid_client_session = Mock(return_value=mock_session_from_transport) mock_connector = Mock(spec=aiohttp.BaseConnector) - handler = BaseLLMAIOHTTPHandler( - transport=mock_transport, connector=mock_connector - ) + handler = BaseLLMAIOHTTPHandler(transport=mock_transport, connector=mock_connector) result = handler._create_client_session_with_transport() # Should use transport, not connector - mock_transport._get_valid_client_session.assert_called_once() + mock_transport.get_valid_client_session.assert_called_once() assert result is mock_session_from_transport diff --git a/tests/unit/llms/custom_httpx/test_aiohttp_so_keepalive.py b/tests/unit/llms/custom_httpx/test_aiohttp_so_keepalive.py index 0065bf8f4ef..a9f84cea57a 100644 --- a/tests/unit/llms/custom_httpx/test_aiohttp_so_keepalive.py +++ b/tests/unit/llms/custom_httpx/test_aiohttp_so_keepalive.py @@ -8,14 +8,12 @@ import pytest def _invoke_connector_factory(http_handler_module): """ Drive the lambda factory installed on the transport so TCPConnector is - actually constructed. _create_aiohttp_transport returns a transport whose + actually constructed. create_aiohttp_transport returns a transport whose _client_factory is the lambda that builds (TCPConnector → ClientSession); - invoking it directly avoids relying on _get_valid_client_session's internal + invoking it directly avoids relying on get_valid_client_session's internal branching to trigger connector construction. """ - transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport( - shared_session=None - ) + transport = http_handler_module.AsyncHTTPHandler.create_aiohttp_transport(shared_session=None) transport._client_factory() return transport @@ -93,7 +91,7 @@ def test_socket_factory_sets_keepalive_options(monkeypatch): monkeypatch.setattr(http_handler_module, "AIOHTTP_TCP_KEEPINTVL", 15) monkeypatch.setattr(http_handler_module, "AIOHTTP_TCP_KEEPCNT", 4) - factory = http_handler_module._build_aiohttp_keepalive_socket_factory() + factory = http_handler_module.build_aiohttp_keepalive_socket_factory() assert factory is not None addr_info = (socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP, "", ("", 0)) @@ -137,7 +135,7 @@ def test_socket_factory_uses_tcp_keepalive_when_keepidle_unavailable(monkeypatch monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True) monkeypatch.setattr(http_handler_module, "AIOHTTP_TCP_KEEPIDLE", 60) - factory = http_handler_module._build_aiohttp_keepalive_socket_factory() + factory = http_handler_module.build_aiohttp_keepalive_socket_factory() assert factory is not None fake_socket_module = MagicMock(spec=[]) @@ -167,7 +165,7 @@ def test_socket_factory_uses_tcp_keepalive_when_keepidle_unavailable(monkeypatch @pytest.mark.asyncio async def test_shared_session_transport_rebuilds_with_socket_factory(monkeypatch): """ - The proxy hands _create_aiohttp_transport an already-built shared session. + The proxy hands create_aiohttp_transport an already-built shared session. When that session is rebuilt (closed session, or a session from another event loop) the replacement must still carry the keep-alive socket factory and the configured keepalive timeout, otherwise AIOHTTP_SO_KEEPALIVE stops @@ -179,14 +177,16 @@ async def test_shared_session_transport_rebuilds_with_socket_factory(monkeypatch monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True) shared_session = aiohttp.ClientSession() - transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport(shared_session=shared_session) + transport = http_handler_module.AsyncHTTPHandler.create_aiohttp_transport(shared_session=shared_session) await shared_session.close() rebuilt_session = MagicMock(name="rebuilt_session") - with patch.object(http_handler_module, "TCPConnector", return_value=MagicMock(name="connector")) as mock_tcp_connector: + with patch.object( + http_handler_module, "TCPConnector", return_value=MagicMock(name="connector") + ) as mock_tcp_connector: with patch.object(http_handler_module, "ClientSession", return_value=rebuilt_session): - assert transport._get_valid_client_session() is rebuilt_session + assert transport.get_valid_client_session() is rebuilt_session assert mock_tcp_connector.call_count == 1 assert callable(mock_tcp_connector.call_args.kwargs.get("socket_factory")) diff --git a/tests/unit/llms/custom_httpx/test_aiohttp_transport.py b/tests/unit/llms/custom_httpx/test_aiohttp_transport.py index 67c1341e356..7c44d61451c 100644 --- a/tests/unit/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/unit/llms/custom_httpx/test_aiohttp_transport.py @@ -771,7 +771,7 @@ async def test_closed_shared_session_rebuild_uses_injected_session_factory(): session_factory=session_factory, # type: ignore ) - assert transport._get_valid_client_session() in rebuilt + assert transport.get_valid_client_session() in rebuilt def test_rebuild_without_running_loop_uses_injected_session_factory(): @@ -792,7 +792,7 @@ def test_rebuild_without_running_loop_uses_injected_session_factory(): session_factory=session_factory, # type: ignore ) - assert transport._get_valid_client_session() in rebuilt + assert transport.get_valid_client_session() in rebuilt @pytest.mark.asyncio @@ -811,7 +811,7 @@ async def test_rebuilt_session_becomes_transport_owned(): session_factory=lambda: replacement, ) - assert transport._get_valid_client_session() is replacement + assert transport.get_valid_client_session() is replacement await transport.aclose() @@ -837,7 +837,7 @@ async def test_stale_loop_rebuild_does_not_close_unowned_session(): try: shared_session._loop = other_loop - assert transport._get_valid_client_session() is replacement + assert transport.get_valid_client_session() is replacement shared_session._loop = running_loop await asyncio.sleep(0.05) assert not shared_session.closed @@ -883,7 +883,7 @@ def _flaky_get_running_loop_factory(): """get_running_loop stand-in that fails once, then delegates. Reproduces #24230: a transient loop-inspection failure sends - _get_valid_client_session into its (RuntimeError, AttributeError) + get_valid_client_session into its (RuntimeError, AttributeError) fallback branch. """ real_get_running_loop = asyncio.get_running_loop @@ -915,7 +915,7 @@ async def test_fallback_recreate_closes_previous_session(): "litellm.llms.custom_httpx.aiohttp_transport.asyncio.get_running_loop", side_effect=_flaky_get_running_loop_factory(), ): - new_session = transport._get_valid_client_session() + new_session = transport.get_valid_client_session() try: assert new_session is not old_session @@ -947,7 +947,7 @@ async def test_replaced_session_emits_no_unclosed_warnings(): "litellm.llms.custom_httpx.aiohttp_transport.asyncio.get_running_loop", side_effect=_flaky_get_running_loop_factory(), ): - new_session = transport._get_valid_client_session() + new_session = transport.get_valid_client_session() try: for _ in range(3): @@ -980,7 +980,7 @@ async def test_dead_loop_session_closed_synchronously_on_recycle(): transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession()) transport.client = old_session - new_session = transport._get_valid_client_session() + new_session = transport.get_valid_client_session() try: assert new_session is not old_session @@ -1040,7 +1040,7 @@ async def test_session_from_other_running_loop_closed_threadsafe(): transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession()) transport.client = holder["session"] - new_session = transport._get_valid_client_session() + new_session = transport.get_valid_client_session() try: deadline = time.monotonic() + 5 @@ -1139,7 +1139,7 @@ async def test_stopped_loop_session_disposed_synchronously_on_recycle(): transport = LiteLLMAiohttpTransport(client=lambda: aiohttp.ClientSession()) transport.client = old_session - new_session = transport._get_valid_client_session() + new_session = transport.get_valid_client_session() try: assert new_session is not old_session @@ -1226,28 +1226,28 @@ def _closed_local_port() -> int: @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") async def test_client_session_helper() -> None: - transport: Final = AsyncHTTPHandler._create_aiohttp_transport() + transport: Final = AsyncHTTPHandler.create_aiohttp_transport() assert isinstance(transport, LiteLLMAiohttpTransport) - session1: Final = transport._get_valid_client_session() + session1: Final = transport.get_valid_client_session() assert isinstance(session1, ClientSession) assert session1.closed is False assert getattr(session1, "_loop") is asyncio.get_running_loop() - session2: Final = transport._get_valid_client_session() + session2: Final = transport.get_valid_client_session() assert session2 is session1 await session1.close() @pytest.mark.usefixtures("_vcr_outcome_gate", "setup_and_teardown") async def test_event_loop_robustness() -> None: - transport: Final = AsyncHTTPHandler._create_aiohttp_transport() - session: Final = transport._get_valid_client_session() + transport: Final = AsyncHTTPHandler.create_aiohttp_transport() + session: Final = transport.get_valid_client_session() assert isinstance(session, ClientSession) await session.close() - session_after_close: Final = transport._get_valid_client_session() + session_after_close: Final = transport.get_valid_client_session() assert isinstance(session_after_close, ClientSession) assert session_after_close is not session assert session_after_close.closed is False transport.client = lambda: ClientSession() - session_after_factory: Final = transport._get_valid_client_session() + session_after_factory: Final = transport.get_valid_client_session() assert isinstance(session_after_factory, ClientSession) assert session_after_factory is not session_after_close assert session_after_factory.closed is False @@ -1261,14 +1261,14 @@ async def test_refused_connection_maps_to_httpx_connect_error( ssl_verify: bool | None, expected_ssl: bool, monkeypatch: pytest.MonkeyPatch ) -> None: monkeypatch.setenv("NO_PROXY", "127.0.0.1") - transport: Final = AsyncHTTPHandler._create_aiohttp_transport(ssl_verify=ssl_verify) + transport: Final = AsyncHTTPHandler.create_aiohttp_transport(ssl_verify=ssl_verify) port: Final = _closed_local_port() request: Final = httpx.Request("GET", f"https://127.0.0.1:{port}/") try: with pytest.raises(httpx.ConnectError) as raised: await transport.handle_async_request(request) finally: - await transport._get_valid_client_session().close() + await transport.get_valid_client_session().close() cause: Final = raised.value.__cause__ assert isinstance(cause, aiohttp_aiohttp_handler.ClientConnectorError) assert cause.ssl is expected_ssl diff --git a/tests/unit/llms/custom_httpx/test_http_handler.py b/tests/unit/llms/custom_httpx/test_http_handler.py index 1d80b0a09dd..f5ae65e9898 100644 --- a/tests/unit/llms/custom_httpx/test_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_http_handler.py @@ -23,7 +23,7 @@ from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, MaskedHTTPStatusError, - _get_httpx_client, + get_httpx_client, get_ssl_configuration, ) from litellm.types.llms.custom_http import VerifyTypes @@ -149,7 +149,7 @@ async def test_ssl_security_level(monkeypatch): assert isinstance(transport, LiteLLMAiohttpTransport) # Get the aiohttp ClientSession - client_session = transport._get_valid_client_session() + client_session = transport.get_valid_client_session() # Get the connector from the session connector = client_session.connector @@ -170,7 +170,7 @@ async def test_force_ipv4_transport(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(litellm, "force_ipv4", True) monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - transport = AsyncHTTPHandler._create_async_transport() + transport = AsyncHTTPHandler.create_async_transport() # Should get an AsyncHTTPTransport (no real HTTP call — avoids CI hangs) assert isinstance(transport, httpx.AsyncHTTPTransport) @@ -182,7 +182,7 @@ async def test_aiohttp_disabled_transport(monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) monkeypatch.setattr(litellm, "force_ipv4", False) - transport = AsyncHTTPHandler._create_async_transport() + transport = AsyncHTTPHandler.create_async_transport() # Should get None when both aiohttp is disabled and force_ipv4 is False assert transport is None @@ -206,7 +206,7 @@ async def test_ssl_verification_with_aiohttp_transport(monkeypatch: pytest.Monke try: transport = litellm_async_client.client._transport assert isinstance(transport, LiteLLMAiohttpTransport) - transport_connector = transport._get_valid_client_session().connector + transport_connector = transport.get_valid_client_session().connector assert isinstance(transport_connector, TCPConnector) aiohttp_session = aiohttp.ClientSession(connector=aiohttp.TCPConnector(ssl=False)) @@ -228,7 +228,7 @@ async def test_ssl_verification_with_shared_session(monkeypatch: pytest.MonkeyPa Test that ssl_verify=False is respected even with shared sessions. This was a bug where shared sessions bypassed SSL configuration because - _create_aiohttp_transport returned immediately without passing ssl_verify + create_aiohttp_transport returned immediately without passing ssl_verify to the LiteLLMAiohttpTransport constructor. The fix stores ssl_verify in the transport and passes it per-request. @@ -242,7 +242,7 @@ async def test_ssl_verification_with_shared_session(monkeypatch: pytest.MonkeyPa try: # Create transport with shared session and ssl_verify=False - transport = AsyncHTTPHandler._create_aiohttp_transport( + transport = AsyncHTTPHandler.create_aiohttp_transport( ssl_verify=False, shared_session=shared_session, ) @@ -273,7 +273,7 @@ async def test_ssl_context_with_shared_session(monkeypatch: pytest.MonkeyPatch): try: # Create transport with shared session and custom ssl_context - transport = AsyncHTTPHandler._create_aiohttp_transport( + transport = AsyncHTTPHandler.create_aiohttp_transport( ssl_context=custom_ssl_context, shared_session=shared_session, ) @@ -337,14 +337,14 @@ class MockClientSession: @pytest.mark.asyncio async def test_create_aiohttp_transport_with_shared_session(): - """Test that _create_aiohttp_transport reuses shared session when provided""" + """Test that create_aiohttp_transport reuses shared session when provided""" from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler # Create a mock shared session that's not callable mock_session = MockClientSession() # Test with shared session - transport = AsyncHTTPHandler._create_aiohttp_transport( + transport = AsyncHTTPHandler.create_aiohttp_transport( shared_session=mock_session # type: ignore ) @@ -420,7 +420,7 @@ async def test_session_reuse_chain(): mock_session = MockClientSession() # Test the entire chain - transport = AsyncHTTPHandler._create_async_transport( + transport = AsyncHTTPHandler.create_async_transport( shared_session=mock_session # type: ignore ) @@ -633,11 +633,11 @@ async def test_httpx_handler_uses_env_user_agent(monkeypatch): def test_get_httpx_client_applies_float_timeout_without_mocking_handler(): """ - Exercise real _get_httpx_client + HTTPHandler: params={'timeout': x} must reach httpx.Client(timeout=...). + Exercise real get_httpx_client + HTTPHandler: params={'timeout': x} must reach httpx.Client(timeout=...). Uses an uncommon timeout value to avoid colliding with other cached clients in-process. """ timeout = 3847.291 - handler = _get_httpx_client(params={"timeout": timeout}) + handler = get_httpx_client(params={"timeout": timeout}) try: assert isinstance(handler, HTTPHandler) assert handler.client.timeout == httpx.Timeout(timeout) @@ -647,7 +647,7 @@ def test_get_httpx_client_applies_float_timeout_without_mocking_handler(): def test_get_httpx_client_applies_httpx_timeout_object_without_mocking_handler(): t = httpx.Timeout(40.0, connect=5.0) - handler = _get_httpx_client(params={"timeout": t}) + handler = get_httpx_client(params={"timeout": t}) try: assert handler.client.timeout == t finally: @@ -710,7 +710,7 @@ async def test_async_get_forwards_per_request_timeout(): class TestDefaultCachedClientTimeoutHonorsRequestTimeout: """Cached default httpx clients must fall back to an explicit litellm.request_timeout. - Regression for LIT-2369: get_async_httpx_client / _get_httpx_client hardcoded a + Regression for LIT-2369: get_async_httpx_client / get_httpx_client hardcoded a 600s default and never consulted litellm.request_timeout, so provider calls with no per-model timeout (e.g. Bedrock) hung for 600s. """ @@ -1162,7 +1162,7 @@ async def test_async_close_leaves_assigned_client_open(): def test_client_handed_out_by_sync_cache_survives_eviction_and_collection(fresh_llm_client_cache): from litellm.caching.llm_caching_handler import LLMClientCache - handler = _get_httpx_client() + handler = get_httpx_client() consumer_client = handler.client handler_ref = weakref.ref(handler) @@ -1336,7 +1336,7 @@ def _mint_session_on_dead_loop(handler: AsyncHTTPHandler) -> ClientSession: loop = asyncio.new_event_loop() async def _create() -> ClientSession: - return transport._get_valid_client_session() + return transport.get_valid_client_session() session = loop.run_until_complete(_create()) loop.close() @@ -1368,7 +1368,7 @@ async def test_finalizer_with_running_loop_schedules_close_and_holds_task_ref(): handler = AsyncHTTPHandler(timeout=61.0) transport = handler.client._transport assert isinstance(transport, LiteLLMAiohttpTransport) - session = transport._get_valid_client_session() + session = transport.get_valid_client_session() assert not session.closed del transport @@ -1391,7 +1391,7 @@ async def test_sync_close_helper_respects_session_ownership(): owned_handler = AsyncHTTPHandler(timeout=61.0) owned_transport = owned_handler.client._transport assert isinstance(owned_transport, LiteLLMAiohttpTransport) - owned_session = owned_transport._get_valid_client_session() + owned_session = owned_transport.get_valid_client_session() baseline = set(LiteLLMAiohttpTransport._background_close_tasks) owned_handler._dispose_wrapped_aiohttp_session() @@ -1863,13 +1863,13 @@ async def test_http2_flag_bypasses_aiohttp_transport(monkeypatch: pytest.MonkeyP monkeypatch.delenv("DISABLE_AIOHTTP_TRANSPORT", raising=False) monkeypatch.setattr(litellm, "http2", True) - assert AsyncHTTPHandler._should_use_aiohttp_transport() is False - assert AsyncHTTPHandler._create_async_transport() is None + assert AsyncHTTPHandler.should_use_aiohttp_transport() is False + assert AsyncHTTPHandler.create_async_transport() is None monkeypatch.setattr(litellm, "http2", False) monkeypatch.setenv("LITELLM_HTTP2", "True") - assert AsyncHTTPHandler._should_use_aiohttp_transport() is False - assert AsyncHTTPHandler._create_async_transport() is None + assert AsyncHTTPHandler.should_use_aiohttp_transport() is False + assert AsyncHTTPHandler.create_async_transport() is None @pytest.mark.asyncio @@ -1879,7 +1879,7 @@ async def test_http2_disabled_by_default(monkeypatch: pytest.MonkeyPatch): monkeypatch.delenv("DISABLE_AIOHTTP_TRANSPORT", raising=False) monkeypatch.setattr(litellm, "disable_aiohttp_transport", False) - assert AsyncHTTPHandler._should_use_aiohttp_transport() is True + assert AsyncHTTPHandler.should_use_aiohttp_transport() is True @pytest.fixture() @@ -2007,7 +2007,7 @@ def test_post_delay_exceeds_per_request_timeout_raises(): threading.Thread(target=server.serve_forever, daemon=True).start() host, port = server.server_address - handler = _get_httpx_client(params={"timeout": _CLIENT_DEFAULT_TIMEOUT_S}) + handler = get_httpx_client(params={"timeout": _CLIENT_DEFAULT_TIMEOUT_S}) try: with pytest.raises(LitellmTimeout): handler.post( diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py index 2d59ab337d4..82cc9e9d3ac 100644 --- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -2979,7 +2979,7 @@ def test_vector_store_search_handler_direct_config_sync_skips_http(): config = _make_stub_direct_vector_store_config(stub_response) logging_obj = Mock() - with patch("litellm.llms.custom_httpx.llm_http_handler._get_httpx_client") as mock_get_client: + with patch("litellm.llms.custom_httpx.llm_http_handler.get_httpx_client") as mock_get_client: result = handler.vector_store_search_handler( vector_store_id="vs_direct", query="q", diff --git a/tests/unit/llms/dashscope/test_dashscope_chat_transformation.py b/tests/unit/llms/dashscope/test_dashscope_chat_transformation.py index 4a394e456f8..bdca3e1c903 100644 --- a/tests/unit/llms/dashscope/test_dashscope_chat_transformation.py +++ b/tests/unit/llms/dashscope/test_dashscope_chat_transformation.py @@ -128,9 +128,7 @@ class TestDashScopeConfig: ] # Call the _transform_messages method directly - transformed_messages = config._transform_messages( - messages=messages, model="qwen-turbo", is_async=False - ) + transformed_messages = config.transform_messages(messages=messages, model="qwen-turbo", is_async=False) # Verify that the content is still in list format and has not been transformed to a string assert isinstance(transformed_messages[0]["content"], list) diff --git a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py index 1cbc9eeb897..4f2892b9c20 100644 --- a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py +++ b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py @@ -265,7 +265,7 @@ def test_transform_messages_sanitizes_empty_content(): {"role": "user", "content": [{"type": "text", "text": ""}]}, {"role": "user", "content": "Hi"}, ] - result = config._transform_messages(messages=messages, model="databricks-claude", is_async=False) + result = config.transform_messages(messages=messages, model="databricks-claude", is_async=False) assert "content" not in result[0] assert result[1]["content"] == "Hi" diff --git a/tests/unit/llms/deepinfra/test_deepinfra_chat_transformation.py b/tests/unit/llms/deepinfra/test_deepinfra_chat_transformation.py index 0865a14fbd9..7ec9216df7b 100644 --- a/tests/unit/llms/deepinfra/test_deepinfra_chat_transformation.py +++ b/tests/unit/llms/deepinfra/test_deepinfra_chat_transformation.py @@ -67,7 +67,7 @@ def test_deepinfra_tool_message_content_transformation(): }, ] - transformed_messages = config._transform_messages( + transformed_messages = config.transform_messages( messages=messages_with_array_content, model="deepinfra/Qwen/Qwen3-235B-A22B" ) @@ -101,7 +101,7 @@ def test_deepinfra_tool_message_content_transformation(): }, ] - transformed_messages_complex = config._transform_messages( + transformed_messages_complex = config.transform_messages( messages=messages_with_complex_content, model="deepinfra/Qwen/Qwen3-235B-A22B" ) @@ -134,7 +134,7 @@ def test_deepinfra_tool_message_content_transformation(): }, ] - transformed_messages_string = config._transform_messages( + transformed_messages_string = config.transform_messages( messages=messages_with_string_content, model="deepinfra/Qwen/Qwen3-235B-A22B" ) @@ -186,7 +186,7 @@ async def test_deepinfra_tool_message_content_transformation_async(): ] # Call with is_async=True - transformed_messages = await config._transform_messages( + transformed_messages = await config.transform_messages( messages=messages_with_array_content, model="deepinfra/Qwen/Qwen3-235B-A22B", is_async=True, diff --git a/tests/unit/llms/deepseek/chat/test_deepseek_chat_transformation.py b/tests/unit/llms/deepseek/chat/test_deepseek_chat_transformation.py index 3f93264d0d0..b6dae4d4f3c 100644 --- a/tests/unit/llms/deepseek/chat/test_deepseek_chat_transformation.py +++ b/tests/unit/llms/deepseek/chat/test_deepseek_chat_transformation.py @@ -151,7 +151,7 @@ class TestDeepSeekVisionMultimodalContent: } def test_user_image_list_forwarded_on_vision_model(self): - result = self.config._transform_messages([self._image_message()], model=self.VISION_MODEL) + result = self.config.transform_messages([self._image_message()], model=self.VISION_MODEL) assert isinstance(result[0]["content"], list) assert result[0]["content"][0]["type"] == "text" @@ -159,13 +159,13 @@ class TestDeepSeekVisionMultimodalContent: assert result[0]["content"][1]["image_url"]["url"] == "https://example.com/image.jpg" def test_image_list_collapsed_on_non_vision_model(self): - result = self.config._transform_messages([self._image_message()], model=self.NON_VISION_MODEL) + result = self.config.transform_messages([self._image_message()], model=self.NON_VISION_MODEL) assert result[0]["content"] == "what is in this image?" def test_image_list_collapsed_on_non_user_roles_even_on_vision_model(self): for role in ("assistant", "system"): - result = self.config._transform_messages([self._image_message(role=role)], model=self.VISION_MODEL) + result = self.config.transform_messages([self._image_message(role=role)], model=self.VISION_MODEL) assert result[0]["content"] == "what is in this image?" @@ -180,7 +180,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert result[0]["content"] == "transcribe this" @@ -195,7 +195,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert result[0]["content"] == "what is this" @@ -210,7 +210,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert isinstance(result[0]["content"], str) assert result[0]["content"] == "Hello world" @@ -219,7 +219,7 @@ class TestDeepSeekVisionMultimodalContent: message = self._image_message() message["search_results"] = [{"source": "kb", "content": [{"text": "article body"}]}] - result = self.config._transform_messages([message], model=self.VISION_MODEL) + result = self.config.transform_messages([message], model=self.VISION_MODEL) content = result[0]["content"] assert isinstance(content, list) @@ -236,7 +236,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.NON_VISION_MODEL) + result = self.config.transform_messages(messages, model=self.NON_VISION_MODEL) assert result[0]["content"] == "context: kbarticle body" @@ -251,7 +251,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert result[0]["content"] == "what is this?" @@ -263,7 +263,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert result[0]["content"] == "hi" @@ -276,7 +276,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert result[0]["content"] == "hi" @@ -291,7 +291,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) content = result[0]["content"] assert isinstance(content, list) @@ -309,7 +309,7 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert result[0]["content"] == "hi" @@ -323,21 +323,21 @@ class TestDeepSeekVisionMultimodalContent: } ] - result = self.config._transform_messages(messages, model=self.NON_VISION_MODEL) + result = self.config.transform_messages(messages, model=self.NON_VISION_MODEL) assert result[0]["content"] == "summarize the docskbarticle body" def test_plain_string_content_message_unchanged(self): messages = [{"role": "user", "content": "hello"}] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert result[0] is messages[0] def test_empty_content_list_untouched(self): messages = [{"role": "user", "content": []}] - result = self.config._transform_messages(messages, model=self.NON_VISION_MODEL) + result = self.config.transform_messages(messages, model=self.NON_VISION_MODEL) assert result[0]["content"] == [] @@ -354,7 +354,7 @@ class TestDeepSeekVisionMultimodalContent: self._image_message(), ] - result = self.config._transform_messages(messages, model=self.VISION_MODEL) + result = self.config.transform_messages(messages, model=self.VISION_MODEL) assert isinstance(result[0]["content"], list) assert result[1]["content"] == "and then?" diff --git a/tests/unit/llms/featherless_ai/chat/test_featherless_chat_transformation.py b/tests/unit/llms/featherless_ai/chat/test_featherless_chat_transformation.py index 560eaf4f06b..69837425396 100644 --- a/tests/unit/llms/featherless_ai/chat/test_featherless_chat_transformation.py +++ b/tests/unit/llms/featherless_ai/chat/test_featherless_chat_transformation.py @@ -153,9 +153,7 @@ class TestFeatherlessAIConfig: ): monkeypatch.delenv(key, raising=False) monkeypatch.setenv("FEATHERLESS_AI_API_KEY", "key-from-ai-env") - api_base, api_key = config._get_openai_compatible_provider_info( - api_base=None, api_key=None - ) + api_base, api_key = config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert api_key == "key-from-ai-env" assert api_base == "https://api.featherless.ai/v1" @@ -170,15 +168,11 @@ class TestFeatherlessAIConfig: ): monkeypatch.delenv(key, raising=False) monkeypatch.setenv("FEATHERLESS_API_KEY", "key-from-legacy-env") - api_base, api_key = config._get_openai_compatible_provider_info( - api_base=None, api_key=None - ) + api_base, api_key = config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert api_key == "key-from-legacy-env" assert api_base == "https://api.featherless.ai/v1" - def test_get_provider_info_prefers_featherless_ai_key_over_legacy( - self, monkeypatch - ): + def test_get_provider_info_prefers_featherless_ai_key_over_legacy(self, monkeypatch): """Test that FEATHERLESS_AI_API_KEY takes precedence over FEATHERLESS_API_KEY""" config = FeatherlessAIConfig() for key in ( @@ -190,9 +184,7 @@ class TestFeatherlessAIConfig: monkeypatch.delenv(key, raising=False) monkeypatch.setenv("FEATHERLESS_AI_API_KEY", "preferred-key") monkeypatch.setenv("FEATHERLESS_API_KEY", "legacy-key") - _, api_key = config._get_openai_compatible_provider_info( - api_base=None, api_key=None - ) + _, api_key = config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert api_key == "preferred-key" def test_default_api_base(self): diff --git a/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py b/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py index c5171da7947..735d0b1125f 100644 --- a/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py +++ b/tests/unit/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py @@ -27,7 +27,7 @@ from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager FIREWORKS_RESPONSES_URL: Final = "https://api.fireworks.ai/inference/v1/responses" -HTTPX_CLIENT_FACTORY: Final = "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client" +HTTPX_CLIENT_FACTORY: Final = "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client" NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({}) NO_PARAMS: Final[Mapping[str, object]] = MappingProxyType({}) diff --git a/tests/unit/llms/gemini/test_gemini_tts.py b/tests/unit/llms/gemini/test_gemini_tts.py index 4893825373a..a7522f6b315 100644 --- a/tests/unit/llms/gemini/test_gemini_tts.py +++ b/tests/unit/llms/gemini/test_gemini_tts.py @@ -92,17 +92,12 @@ class TestGeminiTTSTransformation: assert "speechConfig" in result assert result["speechConfig"]["languageCode"] == "en-US" - assert ( - result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] - == "Kore" - ) + assert result["speechConfig"]["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore" def test_map_audio_params_language_code(self): config = GoogleAIStudioGeminiConfig() - result = config._map_audio_params( - {"voice": "Kore", "format": "pcm16", "language_code": "de-DE"} - ) + result = config.map_audio_params({"voice": "Kore", "format": "pcm16", "language_code": "de-DE"}) assert result["languageCode"] == "de-DE" assert result["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore" @@ -110,7 +105,7 @@ class TestGeminiTTSTransformation: def test_map_audio_params_no_language_code(self): config = GoogleAIStudioGeminiConfig() - result = config._map_audio_params({"voice": "Kore", "format": "pcm16"}) + result = config.map_audio_params({"voice": "Kore", "format": "pcm16"}) assert "languageCode" not in result assert result["voiceConfig"]["prebuiltVoiceConfig"]["voiceName"] == "Kore" @@ -257,26 +252,22 @@ class TestGeminiTTSSpeechConfigInRequestBody: ("gemini-2.5-pro-tts", "vertex_ai"), ], ) - def test_speechconfig_in_generation_config_transform_request_body( - self, model, custom_llm_provider - ): - """Test that speechConfig is included in generationConfig after _transform_request_body()""" + def test_speechconfig_in_generation_config_transform_request_body(self, model, custom_llm_provider): + """Test that speechConfig is included in generationConfig after transform_request_body()""" from litellm.llms.vertex_ai.gemini.transformation import ( - _transform_request_body, + transform_request_body, ) # Simulate optional_params after map_openai_params() has run optional_params = { - "speechConfig": { - "voiceConfig": {"prebuiltVoiceConfig": {"voiceName": "Kore"}} - }, + "speechConfig": {"voiceConfig": {"prebuiltVoiceConfig": {"voiceName": "Kore"}}}, "responseModalities": ["AUDIO"], } messages = [{"role": "user", "content": "Say hello"}] - # Call _transform_request_body which applies the filtering - request_body = _transform_request_body( + # Call transform_request_body which applies the filtering + request_body = transform_request_body( messages=messages, model=model, optional_params=optional_params, @@ -308,13 +299,13 @@ class TestGeminiTTSSpeechConfigInRequestBody: ], ) def test_speechconfig_end_to_end_mapping(self, model, custom_llm_provider): - """Test full pipeline: audio param -> map_openai_params -> _transform_request_body""" + """Test full pipeline: audio param -> map_openai_params -> transform_request_body""" + from litellm.llms.vertex_ai.gemini.transformation import ( + transform_request_body, + ) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) - from litellm.llms.vertex_ai.gemini.transformation import ( - _transform_request_body, - ) config = VertexGeminiConfig() @@ -335,7 +326,7 @@ class TestGeminiTTSSpeechConfigInRequestBody: messages = [{"role": "user", "content": "Hello world"}] # Step 2: Transform to request body (this is where the bug was) - request_body = _transform_request_body( + request_body = transform_request_body( messages=messages, model=model, optional_params=mapped_params, @@ -348,7 +339,7 @@ class TestGeminiTTSSpeechConfigInRequestBody: assert "generationConfig" in request_body generation_config = request_body["generationConfig"] assert "speechConfig" in generation_config, ( - f"speechConfig was filtered out during _transform_request_body() for model={model}, provider={custom_llm_provider}. " + f"speechConfig was filtered out during transform_request_body() for model={model}, provider={custom_llm_provider}. " "This breaks Gemini TTS - speechConfig must be in GenerationConfig TypedDict." ) assert ( @@ -372,18 +363,16 @@ class TestGeminiTTSSpeechConfigInRequestBody: ], ) def test_language_code_end_to_end_mapping(self, model, custom_llm_provider): + from litellm.llms.vertex_ai.gemini.transformation import ( + transform_request_body, + ) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) - from litellm.llms.vertex_ai.gemini.transformation import ( - _transform_request_body, - ) config = VertexGeminiConfig() - non_default_params = { - "audio": {"voice": "Puck", "format": "pcm16", "language_code": "pt-BR"} - } + non_default_params = {"audio": {"voice": "Puck", "format": "pcm16", "language_code": "pt-BR"}} optional_params = {} mapped_params = config.map_openai_params( @@ -395,7 +384,7 @@ class TestGeminiTTSSpeechConfigInRequestBody: assert mapped_params["speechConfig"]["languageCode"] == "pt-BR" - request_body = _transform_request_body( + request_body = transform_request_body( messages=[{"role": "user", "content": "Hello world"}], model=model, optional_params=mapped_params, diff --git a/tests/unit/llms/gigachat/test_file_handler.py b/tests/unit/llms/gigachat/test_file_handler.py index de83b2ddf5f..0739842d7bf 100644 --- a/tests/unit/llms/gigachat/test_file_handler.py +++ b/tests/unit/llms/gigachat/test_file_handler.py @@ -128,7 +128,7 @@ class TestParseDataUrl: class TestDownloadImageSync: - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") def test_downloads_image_successfully(self, mock_http_handler_cls): mock_client = MagicMock() mock_response = MagicMock() @@ -144,7 +144,7 @@ class TestDownloadImageSync: assert ext == "jpeg" mock_client.get.assert_called_once_with("https://example.com/img.jpg") - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") def test_raises_on_http_error(self, mock_http_handler_cls): mock_client = MagicMock() mock_client.get.side_effect = httpx.HTTPStatusError( @@ -157,7 +157,7 @@ class TestDownloadImageSync: with pytest.raises(httpx.HTTPStatusError): file_handler._download_image_sync("https://example.com/404") - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") def test_parse_content_type_fallback(self, mock_http_handler_cls): mock_client = MagicMock() mock_response = MagicMock() @@ -171,7 +171,7 @@ class TestDownloadImageSync: assert content_type == "image/jpeg" assert ext == "jpeg" - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") def test_extracts_extension_from_parametrized_type(self, mock_http_handler_cls): mock_client = MagicMock() mock_response = MagicMock() @@ -235,7 +235,7 @@ class TestDownloadImageAsync: class TestUploadFileSync: @patch(f"{FILE_MODULE}.get_api_base", return_value="https://api.example.com") @patch(f"{FILE_MODULE}.get_access_token", return_value="test-token") - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") def test_uploads_base64_image_and_caches( self, mock_http_handler_cls, mock_get_token, mock_get_api_base ): @@ -268,7 +268,7 @@ class TestUploadFileSync: @patch(f"{FILE_MODULE}.get_api_base", return_value="https://api.example.com") @patch(f"{FILE_MODULE}.get_access_token", return_value="test-token") - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") def test_returns_cached_file_id( self, mock_http_handler_cls, mock_get_token, mock_get_api_base ): @@ -282,7 +282,7 @@ class TestUploadFileSync: # No upload call was made mock_http_handler_cls.return_value.post.assert_not_called() - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") @patch(f"{FILE_MODULE}.get_access_token", return_value="test-token") @patch(f"{FILE_MODULE}.get_api_base", return_value="https://api.example.com") @patch(f"{FILE_MODULE}._download_image_sync") @@ -304,7 +304,7 @@ class TestUploadFileSync: assert result == "file-remote" mock_download.assert_called_once_with("https://example.com/remote.png") - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") @patch(f"{FILE_MODULE}.get_access_token", return_value="test-token") @patch(f"{FILE_MODULE}.get_api_base", return_value="https://api.example.com") def test_returns_none_on_upload_failure( @@ -325,7 +325,7 @@ class TestUploadFileSync: assert result is None - @patch(f"{FILE_MODULE}._get_httpx_client") + @patch(f"{FILE_MODULE}.get_httpx_client") @patch(f"{FILE_MODULE}.get_access_token", return_value="test-token") @patch(f"{FILE_MODULE}.get_api_base", return_value="https://api.example.com") def test_returns_none_when_response_missing_id( diff --git a/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py b/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py index a49a4b44b74..21d4901cf81 100644 --- a/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py +++ b/tests/unit/llms/github_copilot/test_github_copilot_authenticator.py @@ -172,7 +172,7 @@ class TestGitHubCopilotAuthenticator: with ( patch.object(authenticator, "get_access_token", return_value=mock_token), patch( - "litellm.llms.github_copilot.authenticator._get_httpx_client", + "litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client, ), patch.object(mock_response, "json", return_value=mock_api_key_data), @@ -190,7 +190,7 @@ class TestGitHubCopilotAuthenticator: with ( patch.object(authenticator, "get_access_token", return_value=mock_token), patch( - "litellm.llms.github_copilot.authenticator._get_httpx_client", + "litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client, ), patch.object(mock_response, "json", return_value={}), @@ -210,7 +210,7 @@ class TestGitHubCopilotAuthenticator: with ( patch( - "litellm.llms.github_copilot.authenticator._get_httpx_client", + "litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client, ), patch.object(mock_response, "json", return_value=mock_device_code_data), @@ -226,7 +226,7 @@ class TestGitHubCopilotAuthenticator: with ( patch( - "litellm.llms.github_copilot.authenticator._get_httpx_client", + "litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client, ), patch.object(mock_response, "json", return_value=mock_token_data), @@ -284,8 +284,10 @@ class TestGitHubCopilotAuthenticator: "user_code": "UC", "verification_uri": "https://example.com", } - with patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}), \ - patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client): + with ( + patch.dict(os.environ, {"GITHUB_COPILOT_DEVICE_CODE_URL": custom_url}), + patch("litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client), + ): authenticator._get_device_code() assert mock_client.post.call_args[0][0] == custom_url @@ -298,8 +300,10 @@ class TestGitHubCopilotAuthenticator: "user_code": "UC", "verification_uri": "https://example.com", } - with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \ - patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client): + with ( + patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), + patch("litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client), + ): authenticator._get_device_code() assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id @@ -308,9 +312,11 @@ class TestGitHubCopilotAuthenticator: mock_client, mock_response = mock_http_client custom_url = "https://custom.example.com/token" mock_response.json.return_value = {"access_token": "tok"} - with patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}), \ - patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \ - patch("time.sleep"): + with ( + patch.dict(os.environ, {"GITHUB_COPILOT_ACCESS_TOKEN_URL": custom_url}), + patch("litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client), + patch("time.sleep"), + ): authenticator._poll_for_access_token("dc") assert mock_client.post.call_args[0][0] == custom_url @@ -319,9 +325,11 @@ class TestGitHubCopilotAuthenticator: mock_client, mock_response = mock_http_client custom_id = "custom_client_id" mock_response.json.return_value = {"access_token": "tok"} - with patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), \ - patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \ - patch("time.sleep"): + with ( + patch.dict(os.environ, {"GITHUB_COPILOT_CLIENT_ID": custom_id}), + patch("litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client), + patch("time.sleep"), + ): authenticator._poll_for_access_token("dc") assert mock_client.post.call_args[1]["json"]["client_id"] == custom_id @@ -330,9 +338,10 @@ class TestGitHubCopilotAuthenticator: mock_client, mock_response = mock_http_client custom_url = "https://custom.example.com/api-key" mock_response.json.return_value = {"token": "api-tok", "expires_at": 9999999999} - with patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}), \ - patch("litellm.llms.github_copilot.authenticator._get_httpx_client", return_value=mock_client), \ - patch.object(authenticator, "get_access_token", return_value="access-tok"): + with ( + patch.dict(os.environ, {"GITHUB_COPILOT_API_KEY_URL": custom_url}), + patch("litellm.llms.github_copilot.authenticator.get_httpx_client", return_value=mock_client), + patch.object(authenticator, "get_access_token", return_value="access-tok"), + ): authenticator._refresh_api_key() assert mock_client.get.call_args[0][0] == custom_url - diff --git a/tests/unit/llms/github_copilot/test_github_copilot_transformation.py b/tests/unit/llms/github_copilot/test_github_copilot_transformation.py index b1d7e49fbac..1715d106e0a 100644 --- a/tests/unit/llms/github_copilot/test_github_copilot_transformation.py +++ b/tests/unit/llms/github_copilot/test_github_copilot_transformation.py @@ -47,7 +47,7 @@ def test_github_copilot_config_get_openai_compatible_provider_info(): api_base, dynamic_api_key, custom_llm_provider, - ) = config._get_openai_compatible_provider_info( + ) = config.get_openai_compatible_provider_info( model=model, api_base=None, api_key=None, @@ -64,7 +64,7 @@ def test_github_copilot_config_get_openai_compatible_provider_info(): api_base, dynamic_api_key, custom_llm_provider, - ) = config._get_openai_compatible_provider_info( + ) = config.get_openai_compatible_provider_info( model=model, api_base=None, api_key=None, @@ -79,7 +79,7 @@ def test_github_copilot_config_get_openai_compatible_provider_info(): ) with pytest.raises(AuthenticationError) as excinfo: - config._get_openai_compatible_provider_info( + config.get_openai_compatible_provider_info( model=model, api_base=None, api_key=None, @@ -158,25 +158,19 @@ def test_transform_messages_disable_copilot_system_to_assistant(monkeypatch): {"role": "system", "content": "System message."}, {"role": "user", "content": "User message."}, ] - out = config._transform_messages( - [m.copy() for m in messages], model="github_copilot/gpt-4" - ) + out = config.transform_messages([m.copy() for m in messages], model="github_copilot/gpt-4") assert out[0]["role"] == "assistant" assert out[1]["role"] == "user" # Case 2: Flag is True (conversion does not happen) litellm.disable_copilot_system_to_assistant = True - out = config._transform_messages( - [m.copy() for m in messages], model="github_copilot/gpt-4" - ) + out = config.transform_messages([m.copy() for m in messages], model="github_copilot/gpt-4") assert out[0]["role"] == "system" assert out[1]["role"] == "user" # Case 3: Flag is False again (conversion happens) litellm.disable_copilot_system_to_assistant = False - out = config._transform_messages( - [m.copy() for m in messages], model="github_copilot/gpt-4" - ) + out = config.transform_messages([m.copy() for m in messages], model="github_copilot/gpt-4") assert out[0]["role"] == "assistant" assert out[1]["role"] == "user" finally: diff --git a/tests/unit/llms/gradient_ai/chat/test_gradient_ai_chat_transformation.py b/tests/unit/llms/gradient_ai/chat/test_gradient_ai_chat_transformation.py index 6586f970b80..10a7e2b33fa 100644 --- a/tests/unit/llms/gradient_ai/chat/test_gradient_ai_chat_transformation.py +++ b/tests/unit/llms/gradient_ai/chat/test_gradient_ai_chat_transformation.py @@ -74,7 +74,7 @@ def test_transform_messages_handles_dicts_only(config): {"role": "assistant", "content": "Hello!"}, {"role": "user", "content": "Hi!"}, ] - out = config._transform_messages(messages, model="gradient_ai/test-model") + out = config.transform_messages(messages, model="gradient_ai/test-model") assert out[0]["role"] == "assistant" assert out[0]["content"] == "Hello!" assert out[1]["role"] == "user" @@ -84,7 +84,7 @@ def test_transform_messages_handles_dicts_only(config): def test_get_openai_compatible_provider_info_env(monkeypatch, config): monkeypatch.setenv("GRADIENT_AI_AGENT_ENDPOINT", DO_BASE_URL) monkeypatch.setenv("GRADIENT_AI_API_KEY", "env-key") - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == DO_BASE_URL assert api_key == "env-key" @@ -92,6 +92,6 @@ def test_get_openai_compatible_provider_info_env(monkeypatch, config): def test_get_openai_compatible_provider_info_default(monkeypatch, config): monkeypatch.delenv("GRADIENT_AI_AGENT_ENDPOINT", raising=False) monkeypatch.setenv("GRADIENT_AI_API_KEY", "env-key") - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == GRADIENT_AI_SERVERLESS_ENDPOINT assert api_key == "env-key" diff --git a/tests/unit/llms/hosted_vllm/responses/test_hosted_vllm_responses.py b/tests/unit/llms/hosted_vllm/responses/test_hosted_vllm_responses.py index 55d0ce1e68e..e32f876a6b9 100644 --- a/tests/unit/llms/hosted_vllm/responses/test_hosted_vllm_responses.py +++ b/tests/unit/llms/hosted_vllm/responses/test_hosted_vllm_responses.py @@ -71,7 +71,7 @@ def test_hosted_vllm_responses_create_with_string_input(): mock_client = _make_mock_http_client(_make_mock_responses_api_response("I'm doing well, thanks!")) with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ): response = litellm.responses( diff --git a/tests/unit/llms/hyperbolic/chat/test_hyperbolic_transformation.py b/tests/unit/llms/hyperbolic/chat/test_hyperbolic_transformation.py index 0236dc4cd6c..0bc538ffae8 100644 --- a/tests/unit/llms/hyperbolic/chat/test_hyperbolic_transformation.py +++ b/tests/unit/llms/hyperbolic/chat/test_hyperbolic_transformation.py @@ -42,13 +42,13 @@ def test_hyperbolic_get_openai_compatible_provider_info(): config = HyperbolicChatConfig() # Test default API base - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://api.hyperbolic.xyz/v1" # api_key may be set from environment, so we don't test for None # Test custom API base custom_base = "https://custom.hyperbolic.com/v1" - api_base, api_key = config._get_openai_compatible_provider_info(custom_base, "test-key") + api_base, api_key = config.get_openai_compatible_provider_info(custom_base, "test-key") assert api_base == custom_base assert api_key == "test-key" diff --git a/tests/unit/llms/inception/test_inception_chat_transformation.py b/tests/unit/llms/inception/test_inception_chat_transformation.py index c4c023077fc..85d9bcea469 100644 --- a/tests/unit/llms/inception/test_inception_chat_transformation.py +++ b/tests/unit/llms/inception/test_inception_chat_transformation.py @@ -143,7 +143,7 @@ def test_inception_get_openai_compatible_provider_info(): with mock.patch.dict(os.environ, {}, clear=True): with mock.patch.object(litellm, "inception_key", None): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://api.inceptionlabs.ai/v1" assert api_key is None @@ -154,7 +154,7 @@ def test_inception_get_openai_compatible_provider_info(): "INCEPTION_API_BASE": "https://custom.inceptionlabs.ai/v1", }, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://custom.inceptionlabs.ai/v1" assert api_key == "test-key" @@ -165,9 +165,7 @@ def test_inception_get_openai_compatible_provider_info(): "INCEPTION_API_BASE": "https://env.inceptionlabs.ai/v1", }, ): - api_base, api_key = config._get_openai_compatible_provider_info( - "https://param.inceptionlabs.ai/v1", "param-key" - ) + api_base, api_key = config.get_openai_compatible_provider_info("https://param.inceptionlabs.ai/v1", "param-key") assert api_base == "https://param.inceptionlabs.ai/v1" assert api_key == "param-key" @@ -177,7 +175,7 @@ def test_inception_key_module_attr_fallback(): config = InceptionChatConfig() with mock.patch.dict(os.environ, {}, clear=True): with mock.patch.object(litellm, "inception_key", "module-attr-key"): - _, api_key = config._get_openai_compatible_provider_info(None, None) + _, api_key = config.get_openai_compatible_provider_info(None, None) assert api_key == "module-attr-key" @@ -191,16 +189,16 @@ def test_inception_does_not_leak_key_to_caller_api_base(): with mock.patch.dict(os.environ, {"INCEPTION_API_KEY": "server-secret"}, clear=True): with mock.patch.object(litellm, "inception_key", "module-secret"): # caller overrides api_base without a key -> server key withheld - api_base, api_key = config._get_openai_compatible_provider_info("https://attacker.example/v1", None) + api_base, api_key = config.get_openai_compatible_provider_info("https://attacker.example/v1", None) assert api_base == "https://attacker.example/v1" assert api_key is None # caller overrides api_base AND supplies their own key -> used as-is - _, api_key = config._get_openai_compatible_provider_info("https://attacker.example/v1", "caller-key") + _, api_key = config.get_openai_compatible_provider_info("https://attacker.example/v1", "caller-key") assert api_key == "caller-key" # default/server base -> server-managed key resolved - _, api_key = config._get_openai_compatible_provider_info(None, None) + _, api_key = config.get_openai_compatible_provider_info(None, None) assert api_key == "module-secret" diff --git a/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py b/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py index 44e8145ff33..579ba2d75ef 100644 --- a/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py +++ b/tests/unit/llms/jina_ai/embedding/test_jina_embedding_transformation.py @@ -81,7 +81,7 @@ class TestJinaAIEmbeddingTransform: sentinel = f"resolved-via-{env_name.lower()}" monkeypatch.setenv(env_name, sentinel) - _, _, dynamic_api_key = self.config._get_openai_compatible_provider_info(api_base=None, api_key=None) + _, _, dynamic_api_key = self.config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert dynamic_api_key == sentinel @@ -95,7 +95,7 @@ class TestJinaAIEmbeddingTransform: monkeypatch.setenv(name, f"resolved-via-{name.lower()}") for expected_name in JINA_KEY_ENV_NAMES: - _, _, dynamic_api_key = self.config._get_openai_compatible_provider_info(api_base=None, api_key=None) + _, _, dynamic_api_key = self.config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert dynamic_api_key == f"resolved-via-{expected_name.lower()}" monkeypatch.delenv(expected_name) @@ -108,7 +108,7 @@ class TestJinaAIEmbeddingTransform: for name in JINA_KEY_ENV_NAMES: monkeypatch.setenv(name, f"resolved-via-{name.lower()}") - _, _, dynamic_api_key = self.config._get_openai_compatible_provider_info( + _, _, dynamic_api_key = self.config.get_openai_compatible_provider_info( api_base=None, api_key="passed-in-by-caller" ) diff --git a/tests/unit/llms/lemonade/test_lemonade.py b/tests/unit/llms/lemonade/test_lemonade.py index fa0d9d279a7..b17fc039361 100644 --- a/tests/unit/llms/lemonade/test_lemonade.py +++ b/tests/unit/llms/lemonade/test_lemonade.py @@ -27,9 +27,7 @@ def test_get_openai_compatible_provider_info(monkeypatch): monkeypatch.setattr(litellm, "api_key", None) config = LemonadeChatConfig() - api_base, key = config._get_openai_compatible_provider_info( - api_base=None, api_key=None - ) + api_base, key = config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert api_base == "http://localhost:8000/api/v1" assert key == "lemonade" @@ -43,9 +41,7 @@ def test_get_openai_compatible_provider_info_with_custom_base(monkeypatch): config = LemonadeChatConfig() custom_api_base = "https://custom.lemonade.ai/v1" - api_base, key = config._get_openai_compatible_provider_info( - api_base=custom_api_base, api_key=None - ) + api_base, key = config.get_openai_compatible_provider_info(api_base=custom_api_base, api_key=None) assert api_base == custom_api_base assert key == "lemonade" @@ -58,9 +54,7 @@ def test_get_openai_compatible_provider_info_with_api_key_env(monkeypatch): monkeypatch.setattr(litellm, "api_key", None) config = LemonadeChatConfig() - api_base, key = config._get_openai_compatible_provider_info( - api_base=None, api_key=None - ) + api_base, key = config.get_openai_compatible_provider_info(api_base=None, api_key=None) assert api_base == "http://localhost:8000/api/v1" assert key == "test-key" @@ -75,9 +69,7 @@ def test_get_openai_compatible_provider_info_skips_env_key_for_custom_base( monkeypatch.setattr(litellm, "api_key", None) config = LemonadeChatConfig() - api_base, key = config._get_openai_compatible_provider_info( - api_base="https://attacker.example/v1", api_key=None - ) + api_base, key = config.get_openai_compatible_provider_info(api_base="https://attacker.example/v1", api_key=None) assert api_base == "https://attacker.example/v1" assert key == "lemonade" @@ -93,15 +85,13 @@ def test_get_openai_compatible_provider_info_uses_explicit_key_for_custom_base( monkeypatch.setattr(litellm, "api_key", None) config = LemonadeChatConfig() - api_base, key = config._get_openai_compatible_provider_info( + api_base, key = config.get_openai_compatible_provider_info( api_base="https://lemonade.example/v1", api_key="explicit-lemonade-key" ) assert api_base == "https://lemonade.example/v1" assert key == "explicit-lemonade-key" - assert config._get_auth_headers(key) == { - "Authorization": "Bearer explicit-lemonade-key" - } + assert config._get_auth_headers(key) == {"Authorization": "Bearer explicit-lemonade-key"} def test_get_openai_compatible_provider_info_empty_key_does_not_leak_to_custom_base( @@ -113,9 +103,7 @@ def test_get_openai_compatible_provider_info_empty_key_does_not_leak_to_custom_b monkeypatch.setattr(litellm, "api_key", None) config = LemonadeChatConfig() - api_base, key = config._get_openai_compatible_provider_info( - api_base="https://attacker.example/v1", api_key="" - ) + api_base, key = config.get_openai_compatible_provider_info(api_base="https://attacker.example/v1", api_key="") assert api_base == "https://attacker.example/v1" assert key == "lemonade" @@ -129,9 +117,7 @@ def test_get_openai_compatible_provider_info_ignores_global_api_key(monkeypatch) monkeypatch.setattr(litellm, "api_key", "global-openai-key") config = LemonadeChatConfig() - api_base, key = config._get_openai_compatible_provider_info( - api_base="http://lemonade.test/v1", api_key=None - ) + api_base, key = config.get_openai_compatible_provider_info(api_base="http://lemonade.test/v1", api_key=None) assert api_base == "http://lemonade.test/v1" assert key == "lemonade" @@ -393,9 +379,7 @@ def test_transform_response(): model_response = ModelResponse() # Mock the parent class transform_response method - with patch.object( - config.__class__.__bases__[0], "transform_response" - ) as mock_parent: + with patch.object(config.__class__.__bases__[0], "transform_response") as mock_parent: mock_parent.return_value = model_response result = config.transform_response( diff --git a/tests/unit/llms/llamafile/chat/test_llamafile_chat_transformation.py b/tests/unit/llms/llamafile/chat/test_llamafile_chat_transformation.py index 6752098901c..e55938ac2e2 100644 --- a/tests/unit/llms/llamafile/chat/test_llamafile_chat_transformation.py +++ b/tests/unit/llms/llamafile/chat/test_llamafile_chat_transformation.py @@ -135,9 +135,7 @@ def test_get_openai_compatible_provider_info( patch_base as mock_base, patch_key as mock_key, ): - result_base, result_key = config._get_openai_compatible_provider_info( - api_base, api_key - ) + result_base, result_key = config.get_openai_compatible_provider_info(api_base, api_key) assert result_base == expected_base assert result_key == expected_key diff --git a/tests/unit/llms/lm_studio/test_lm_studio_chat_transformation.py b/tests/unit/llms/lm_studio/test_lm_studio_chat_transformation.py index 9a4af91b736..a5d77bbee8e 100644 --- a/tests/unit/llms/lm_studio/test_lm_studio_chat_transformation.py +++ b/tests/unit/llms/lm_studio/test_lm_studio_chat_transformation.py @@ -61,11 +61,11 @@ def test_lm_studio_get_openai_compatible_provider_info(): config = LMStudioChatConfig() # Test default behavior (no API key provided) - _, api_key = config._get_openai_compatible_provider_info(None, None) + _, api_key = config.get_openai_compatible_provider_info(None, None) assert api_key == "fake-api-key" # Test explicit API key - _, api_key = config._get_openai_compatible_provider_info(None, "test-key") + _, api_key = config.get_openai_compatible_provider_info(None, "test-key") assert api_key == "test-key" @@ -80,6 +80,6 @@ def test_lm_studio_get_openai_compatible_provider_info_with_env(): "LM_STUDIO_API_KEY": "env_api_key", }, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "http://localhost:1234/v1" assert api_key == "env_api_key" diff --git a/tests/unit/llms/meta_llama/test_meta_llama_chat_transformation.py b/tests/unit/llms/meta_llama/test_meta_llama_chat_transformation.py index 15995d873c2..9f094b83824 100644 --- a/tests/unit/llms/meta_llama/test_meta_llama_chat_transformation.py +++ b/tests/unit/llms/meta_llama/test_meta_llama_chat_transformation.py @@ -52,15 +52,15 @@ def test_llama_api_streaming_no_307_error(): from litellm.llms.openai.common_utils import BaseOpenAILLM # Verify the async httpx client has follow_redirects enabled - async_client = BaseOpenAILLM._get_async_http_client() + async_client = BaseOpenAILLM.get_async_http_client() assert async_client is not None - assert ( - async_client.follow_redirects is True - ), "Async httpx client should set follow_redirects=True to prevent 307 errors" + assert async_client.follow_redirects is True, ( + "Async httpx client should set follow_redirects=True to prevent 307 errors" + ) # Verify the sync httpx client has follow_redirects enabled sync_client = BaseOpenAILLM._get_sync_http_client() assert sync_client is not None - assert ( - sync_client.follow_redirects is True - ), "Sync httpx client should set follow_redirects=True to prevent 307 errors" + assert sync_client.follow_redirects is True, ( + "Sync httpx client should set follow_redirects=True to prevent 307 errors" + ) diff --git a/tests/unit/llms/mistral/test_mistral_chat_transformation.py b/tests/unit/llms/mistral/test_mistral_chat_transformation.py index 8fb3b3c43df..23b58c789d2 100644 --- a/tests/unit/llms/mistral/test_mistral_chat_transformation.py +++ b/tests/unit/llms/mistral/test_mistral_chat_transformation.py @@ -4,20 +4,18 @@ from unittest.mock import MagicMock, patch import pytest from litellm.litellm_core_utils.prompt_templates.common_utils import TOOL_RESULT_IMAGE_BOUNDARY -from litellm.types.llms.openai import AllMessageValues - - from litellm.llms.mistral.chat.transformation import ( MistralChatResponseIterator, MistralConfig, ) +from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ModelResponse @pytest.mark.asyncio async def test_mistral_chat_transformation(): mistral_config = MistralConfig() - result = mistral_config._transform_messages( + result = mistral_config.transform_messages( **{ "messages": [ { @@ -816,18 +814,14 @@ class TestMistralStripsOutputOnlyFields: "role": "assistant", "content": "Follow-up", "reasoning_content": "Some internal reasoning text.", - "thinking_blocks": [ - {"type": "thinking", "thinking": "step", "signature": "mistral"} - ], + "thinking_blocks": [{"type": "thinking", "thinking": "step", "signature": "mistral"}], }, ], ) result = cast( List[AllMessageValues], - MistralConfig()._transform_messages( - messages=messages, model="mistral-medium-3-5" - ), + MistralConfig().transform_messages(messages=messages, model="mistral-medium-3-5"), ) assistant_message = result[-1] @@ -844,9 +838,7 @@ class TestMistralStripsOutputOnlyFields: result = cast( List[AllMessageValues], - MistralConfig()._transform_messages( - messages=messages, model="mistral-medium-3-5" - ), + MistralConfig().transform_messages(messages=messages, model="mistral-medium-3-5"), ) assert result[0].get("reasoning_content") == "noise" @@ -881,9 +873,7 @@ class TestMistralStripsOutputOnlyFields: ): result = cast( List[AllMessageValues], - MistralConfig()._transform_messages( - messages=messages, model="mistral-medium-3-5", is_async=False - ), + MistralConfig().transform_messages(messages=messages, model="mistral-medium-3-5", is_async=False), ) assert "reasoning_content" not in result[-1] diff --git a/tests/unit/llms/modelscope/chat/test_modelscope_chat_transformation.py b/tests/unit/llms/modelscope/chat/test_modelscope_chat_transformation.py index 6fe39798f4f..03c2f0aa563 100644 --- a/tests/unit/llms/modelscope/chat/test_modelscope_chat_transformation.py +++ b/tests/unit/llms/modelscope/chat/test_modelscope_chat_transformation.py @@ -106,7 +106,7 @@ class TestModelScopeConfig: } ] - result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + result = config.transform_messages(messages=messages, model=DEFAULT_MODEL) assert result[0]["content"] == "Hello world" @@ -123,7 +123,7 @@ class TestModelScopeConfig: } ] - result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + result = config.transform_messages(messages=messages, model=DEFAULT_MODEL) assert isinstance(result[0]["content"], list) assert len(result[0]["content"]) == 2 @@ -135,7 +135,7 @@ class TestModelScopeConfig: config = ModelScopeChatConfig() messages = [{"role": "user", "content": "Hello"}] - result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + result = config.transform_messages(messages=messages, model=DEFAULT_MODEL) assert result[0]["content"] == "Hello" @@ -153,7 +153,7 @@ class TestModelScopeConfig: }, ] - result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + result = config.transform_messages(messages=messages, model=DEFAULT_MODEL) assert result[0]["content"] == "Hi" assert result[1]["content"] == "Hello!" @@ -174,7 +174,7 @@ class TestModelScopeConfig: }, ] - result = config._transform_messages(messages=messages, model=DEFAULT_MODEL) + result = config.transform_messages(messages=messages, model=DEFAULT_MODEL) assert result[0]["content"] == "Hi" assert result[1]["content"] == "Hello!" @@ -233,7 +233,7 @@ class TestModelScopeConfig: """Explicit api_base and api_key should be returned as-is.""" config = ModelScopeChatConfig() - api_base, api_key = config._get_openai_compatible_provider_info( + api_base, api_key = config.get_openai_compatible_provider_info( api_base="https://custom.example.com/v1", api_key="my-key", ) @@ -249,7 +249,7 @@ class TestModelScopeConfig: os.environ.pop("MODELSCOPE_API_BASE", None) os.environ.pop("MODELSCOPE_API_KEY", None) - api_base, api_key = config._get_openai_compatible_provider_info( + api_base, api_key = config.get_openai_compatible_provider_info( api_base=None, api_key=None, ) @@ -265,7 +265,7 @@ class TestModelScopeConfig: os.environ, {"MODELSCOPE_API_BASE": "https://env.modelscope.cn/v1"}, ): - api_base, _ = config._get_openai_compatible_provider_info( + api_base, _ = config.get_openai_compatible_provider_info( api_base=None, api_key=None, ) diff --git a/tests/unit/llms/neosantara/test_neosantara.py b/tests/unit/llms/neosantara/test_neosantara.py index bef8c60d171..8276a598e4b 100644 --- a/tests/unit/llms/neosantara/test_neosantara.py +++ b/tests/unit/llms/neosantara/test_neosantara.py @@ -34,7 +34,7 @@ def test_neosantara_dynamic_config_env_vars(): "NEOSANTARA_API_BASE": "https://custom.neosantara.example/v1", }, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://custom.neosantara.example/v1" assert api_key == "test-key" diff --git a/tests/unit/llms/oci/chat/test_oci_cohere_tool_calls.py b/tests/unit/llms/oci/chat/test_oci_cohere_tool_calls.py index 9b06b01aa00..30b574e040e 100644 --- a/tests/unit/llms/oci/chat/test_oci_cohere_tool_calls.py +++ b/tests/unit/llms/oci/chat/test_oci_cohere_tool_calls.py @@ -940,16 +940,16 @@ class TestCohereMessageAdaptationEdgeCases: assert history[0].toolCalls[0].parameters == {} def test_extract_text_content_list_with_non_dict_items(self): - from litellm.llms.oci.chat.cohere import _extract_text_content + from litellm.llms.oci.chat.cohere import extract_text_content # List with a non-dict item — should be silently skipped - result = _extract_text_content([{"type": "text", "text": "hello"}, "bad_item"]) + result = extract_text_content([{"type": "text", "text": "hello"}, "bad_item"]) assert result == "hello" def test_extract_text_content_non_string_non_list(self): - from litellm.llms.oci.chat.cohere import _extract_text_content + from litellm.llms.oci.chat.cohere import extract_text_content - result = _extract_text_content(12345) + result = extract_text_content(12345) assert result == "12345" diff --git a/tests/unit/llms/oci/test_oci_coverage_boost.py b/tests/unit/llms/oci/test_oci_coverage_boost.py index 8f7588c5de7..251df5e4cad 100644 --- a/tests/unit/llms/oci/test_oci_coverage_boost.py +++ b/tests/unit/llms/oci/test_oci_coverage_boost.py @@ -12,19 +12,18 @@ All tests are self-contained and require no real OCI credentials or network acce import json from typing import TYPE_CHECKING - -import pytest -from unittest.mock import patch, MagicMock, AsyncMock +from unittest.mock import AsyncMock, MagicMock, patch import httpx +import pytest if TYPE_CHECKING: from litellm.llms.oci.chat.transformation import OCIStreamWrapper from litellm import ModelResponse from litellm.llms.oci.chat.cohere import ( - _extract_text_content, adapt_messages_to_cohere_standard, + extract_text_content, handle_cohere_response, handle_cohere_stream_chunk, ) @@ -407,16 +406,16 @@ def test_handle_generic_stream_chunk_no_message(): # =========================================================================== -# cohere.py — _extract_text_content +# cohere.py — extract_text_content # =========================================================================== def test_extract_text_content_none(): - assert _extract_text_content(None) == "" + assert extract_text_content(None) == "" def test_extract_text_content_string(): - assert _extract_text_content("hello") == "hello" + assert extract_text_content("hello") == "hello" def test_extract_text_content_list(): @@ -424,7 +423,7 @@ def test_extract_text_content_list(): {"type": "text", "text": "foo"}, {"type": "text", "text": "bar"}, ] - assert _extract_text_content(content) == "foobar" + assert extract_text_content(content) == "foobar" def test_extract_text_content_list_skips_non_text(): @@ -432,11 +431,11 @@ def test_extract_text_content_list_skips_non_text(): {"type": "image_url", "url": "https://x.com/img.png"}, {"type": "text", "text": "only this"}, ] - assert _extract_text_content(content) == "only this" + assert extract_text_content(content) == "only this" def test_extract_text_content_non_string_non_list(): - assert _extract_text_content(42) == "42" + assert extract_text_content(42) == "42" # =========================================================================== diff --git a/tests/unit/llms/openai/completion/test_handler.py b/tests/unit/llms/openai/completion/test_handler.py index 892ec80f570..b4b224f0977 100644 --- a/tests/unit/llms/openai/completion/test_handler.py +++ b/tests/unit/llms/openai/completion/test_handler.py @@ -65,7 +65,7 @@ def test_convert_dict_to_text_completion_response(): @pytest.mark.asyncio async def test_acompletion_uses_optimized_http_client(): """ - Test that OpenAITextCompletion.acompletion uses BaseOpenAILLM._get_async_http_client() + Test that OpenAITextCompletion.acompletion uses BaseOpenAILLM.get_async_http_client() instead of litellm.aclient_session directly. Related issue: https://github.com/BerriAI/litellm/issues/17676 diff --git a/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py b/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py index f7a88b5ba63..8f4fd56c7c6 100644 --- a/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/unit/llms/openai/realtime/test_openai_realtime_handler.py @@ -7,15 +7,12 @@ import pytest from litellm.llms.custom_httpx.http_handler import get_shared_realtime_ssl_context - -@pytest.mark.parametrize( - "api_base", ["https://api.openai.com/v1", "https://api.openai.com"] -) +@pytest.mark.parametrize("api_base", ["https://api.openai.com/v1", "https://api.openai.com"]) def test_openai_realtime_handler_url_construction(api_base): from litellm.llms.openai.realtime.handler import OpenAIRealtime handler = OpenAIRealtime() - url = handler._construct_url( + url = handler.construct_url( api_base=api_base, query_params={ "model": "gpt-4o-realtime-preview-2024-10-01", @@ -36,7 +33,7 @@ def test_openai_realtime_handler_url_with_extra_params(): "model": "gpt-4o-realtime-preview-2024-10-01", "intent": "chat", } - url = handler._construct_url(api_base=api_base, query_params=query_params) + url = handler.construct_url(api_base=api_base, query_params=query_params) # Both 'model' and other params should be included in the query string assert url.startswith("wss://api.openai.com/v1/realtime?") assert "model=gpt-4o-realtime-preview-2024-10-01" in url @@ -59,12 +56,8 @@ def test_openai_realtime_handler_model_parameter_inclusion(): api_base = "https://api.openai.com/" # Test with just model parameter - query_params_model_only: RealtimeQueryParams = { - "model": "gpt-4o-mini-realtime-preview" - } - url = handler._construct_url( - api_base=api_base, query_params=query_params_model_only - ) + query_params_model_only: RealtimeQueryParams = {"model": "gpt-4o-mini-realtime-preview"} + url = handler.construct_url(api_base=api_base, query_params=query_params_model_only) # Verify the URL structure assert url.startswith("wss://api.openai.com/v1/realtime?") @@ -75,9 +68,7 @@ def test_openai_realtime_handler_model_parameter_inclusion(): "model": "gpt-4o-mini-realtime-preview", "intent": "chat", } - url_with_extras = handler._construct_url( - api_base=api_base, query_params=query_params_with_extras - ) + url_with_extras = handler.construct_url(api_base=api_base, query_params=query_params_with_extras) # Verify both parameters are included assert url_with_extras.startswith("wss://api.openai.com/v1/realtime?") diff --git a/tests/unit/llms/openai/realtime/test_transcription_sessions.py b/tests/unit/llms/openai/realtime/test_transcription_sessions.py index 54f206d098d..3bb9a4ab1ab 100644 --- a/tests/unit/llms/openai/realtime/test_transcription_sessions.py +++ b/tests/unit/llms/openai/realtime/test_transcription_sessions.py @@ -207,14 +207,14 @@ def test_azure_construct_url_encodes_model_and_api_version(): from litellm.llms.azure.realtime.handler import AzureOpenAIRealtime h = AzureOpenAIRealtime() - url = h._construct_url( + url = h.construct_url( "https://x.openai.azure.com", "deploy&evil=1", "2024-10-01-preview", ) assert "evil=1" not in url.split("?", 1)[1] - url_ga = h._construct_url( + url_ga = h.construct_url( "https://x.openai.azure.com", "deploy&evil=1", None, diff --git a/tests/unit/llms/openai/responses/test_openai_responses_data_residency.py b/tests/unit/llms/openai/responses/test_openai_responses_data_residency.py index ac89428617d..cf9f9706cfe 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_data_residency.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_data_residency.py @@ -74,7 +74,7 @@ def test_responses_eu_api_base_sets_data_residency(): with ( patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ), patch.object(litellm.Logging, "__init__", init_spy), @@ -96,7 +96,7 @@ def test_responses_us_api_base_sets_data_residency(): with ( patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ), patch.object(litellm.Logging, "__init__", init_spy), @@ -118,7 +118,7 @@ def test_responses_global_api_base_leaves_data_residency_none(): with ( patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ), patch.object(litellm.Logging, "__init__", init_spy), diff --git a/tests/unit/llms/openai/test_openai_common_utils.py b/tests/unit/llms/openai/test_openai_common_utils.py index b54ec10ef17..fb605c6da7a 100644 --- a/tests/unit/llms/openai/test_openai_common_utils.py +++ b/tests/unit/llms/openai/test_openai_common_utils.py @@ -197,7 +197,7 @@ def test_evicting_a_client_built_on_the_callers_session_leaves_that_session_open LLMClientCache(evicted_client_closer=closer), ) - wrapper = OpenAIChatCompletion()._get_openai_client( + wrapper = OpenAIChatCompletion().get_openai_client( is_async=True, api_key="sk-not-a-real-key", api_base="https://api.openai.com/v1", @@ -229,7 +229,7 @@ def test_a_client_litellm_built_its_own_http_client_for_is_still_closed(monkeypa LLMClientCache(evicted_client_closer=closer), ) - wrapper = OpenAIChatCompletion()._get_openai_client( + wrapper = OpenAIChatCompletion().get_openai_client( is_async=False, api_key="sk-not-a-real-key", api_base="https://api.openai.com/v1", diff --git a/tests/unit/llms/openai/test_openai_workload_identity.py b/tests/unit/llms/openai/test_openai_workload_identity.py index 74415d45638..7d78a2c5238 100644 --- a/tests/unit/llms/openai/test_openai_workload_identity.py +++ b/tests/unit/llms/openai/test_openai_workload_identity.py @@ -184,19 +184,19 @@ class TestTokenExchange: class TestClientConstruction: def test_sync_client_uses_workload_identity(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: - client: Final = OpenAIChatCompletion()._get_openai_client(is_async=False, api_key=None, api_base=None) + client: Final = OpenAIChatCompletion().get_openai_client(is_async=False, api_key=None, api_base=None) assert isinstance(client, OpenAI) assert client.api_key == "workload-identity-auth" assert client._workload_identity_auth is not None def test_async_client_uses_workload_identity(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: - client: Final = OpenAIChatCompletion()._get_openai_client(is_async=True, api_key=None, api_base=None) + client: Final = OpenAIChatCompletion().get_openai_client(is_async=True, api_key=None, api_base=None) assert isinstance(client, AsyncOpenAI) assert client.api_key == "workload-identity-auth" assert client._workload_identity_auth is not None def test_privatelink_client_uses_workload_identity(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: - client: Final = OpenAIChatCompletion()._get_openai_client( + client: Final = OpenAIChatCompletion().get_openai_client( is_async=False, api_key=None, api_base="https://southcentralus.privatelink.api.openai.com/v1" ) assert isinstance(client, OpenAI) @@ -204,7 +204,7 @@ class TestClientConstruction: assert client._workload_identity_auth is not None def test_static_key_client_unaffected(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: - client: Final = OpenAIChatCompletion()._get_openai_client(is_async=False, api_key="sk-static", api_base=None) + client: Final = OpenAIChatCompletion().get_openai_client(is_async=False, api_key="sk-static", api_base=None) assert isinstance(client, OpenAI) assert client.api_key == "sk-static" assert client._workload_identity_auth is None @@ -246,7 +246,7 @@ class TestClientConstruction: }, ) ) - client = OpenAIChatCompletion()._get_openai_client(is_async=False, api_key=None, api_base=None) + client = OpenAIChatCompletion().get_openai_client(is_async=False, api_key=None, api_base=None) assert isinstance(client, OpenAI) client.chat.completions.create(model="gpt-4o-mini", messages=[{"role": "user", "content": "hi"}]) auth_header: Final = completion_route.calls.last.request.headers["Authorization"] @@ -420,7 +420,7 @@ class TestResolveConfigFromDeployment: class TestDeploymentClientConstruction: def test_sync_client_from_deployment_params(self, deployment_wif: dict[str, str]) -> None: - client: Final = OpenAIChatCompletion()._get_openai_client( + client: Final = OpenAIChatCompletion().get_openai_client( is_async=False, api_key=None, api_base=None, litellm_params=deployment_wif ) assert isinstance(client, OpenAI) @@ -428,7 +428,7 @@ class TestDeploymentClientConstruction: assert client._workload_identity_auth is not None def test_async_client_from_deployment_params(self, deployment_wif: dict[str, str]) -> None: - client: Final = OpenAIChatCompletion()._get_openai_client( + client: Final = OpenAIChatCompletion().get_openai_client( is_async=True, api_key=None, api_base=None, litellm_params=deployment_wif ) assert isinstance(client, AsyncOpenAI) @@ -437,13 +437,13 @@ class TestDeploymentClientConstruction: def test_distinct_deployments_get_distinct_cached_clients(self, deployment_wif: dict[str, str]) -> None: other_deployment: Final = {**deployment_wif, "openai_service_account_id": "user-other"} handler: Final = OpenAIChatCompletion() - first: Final = handler._get_openai_client( + first: Final = handler.get_openai_client( is_async=False, api_key=None, api_base=None, litellm_params=deployment_wif ) - second: Final = handler._get_openai_client( + second: Final = handler.get_openai_client( is_async=False, api_key=None, api_base=None, litellm_params=other_deployment ) - again: Final = handler._get_openai_client( + again: Final = handler.get_openai_client( is_async=False, api_key=None, api_base=None, litellm_params=dict(deployment_wif) ) assert first is not second diff --git a/tests/unit/llms/openai_like/test_abliteration_provider.py b/tests/unit/llms/openai_like/test_abliteration_provider.py index 8b8d443fc44..33ad8abf62c 100644 --- a/tests/unit/llms/openai_like/test_abliteration_provider.py +++ b/tests/unit/llms/openai_like/test_abliteration_provider.py @@ -32,7 +32,7 @@ def test_abliteration_provider_registered(): def test_abliteration_resolves_env_api_key(monkeypatch): config = _get_config() monkeypatch.setenv("ABLITERATION_API_KEY", "test-key") - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == ABLITERATION_BASE_URL assert api_key == "test-key" diff --git a/tests/unit/llms/openai_like/test_assemblyai_provider.py b/tests/unit/llms/openai_like/test_assemblyai_provider.py index 7eee810b271..c325384e051 100644 --- a/tests/unit/llms/openai_like/test_assemblyai_provider.py +++ b/tests/unit/llms/openai_like/test_assemblyai_provider.py @@ -32,7 +32,7 @@ def test_assemblyai_provider_registered(): def test_assemblyai_resolves_env_api_key(monkeypatch): config = _get_config() monkeypatch.setenv("ASSEMBLYAI_API_KEY", "test-key") - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == ASSEMBLYAI_BASE_URL assert api_key == "test-key" diff --git a/tests/unit/llms/openai_like/test_cognition_provider.py b/tests/unit/llms/openai_like/test_cognition_provider.py index 81895d7dc42..f7529830e44 100644 --- a/tests/unit/llms/openai_like/test_cognition_provider.py +++ b/tests/unit/llms/openai_like/test_cognition_provider.py @@ -104,14 +104,12 @@ class TestCognitionProviderIdentity: provider = JSONProviderRegistry.get("cognition") assert provider is not None - api_base, api_key = create_config_class(provider)()._get_openai_compatible_provider_info(None, None) + api_base, api_key = create_config_class(provider)().get_openai_compatible_provider_info(None, None) assert api_base == "https://api.cognition.ai/v1" assert api_key == "sk-cognition-env" class TestCognitionCostTracking: - - def test_supported_endpoints_matrix(self): matrix = json.loads((Path(litellm.__file__).parent / "provider_endpoints_support_backup.json").read_text()) diff --git a/tests/unit/llms/openai_like/test_empiriolabs_provider.py b/tests/unit/llms/openai_like/test_empiriolabs_provider.py index 58f5e47d09e..a8acce5cc9d 100644 --- a/tests/unit/llms/openai_like/test_empiriolabs_provider.py +++ b/tests/unit/llms/openai_like/test_empiriolabs_provider.py @@ -33,7 +33,7 @@ def test_empiriolabs_provider_registered(): def test_empiriolabs_resolves_env_api_key(monkeypatch): config = _get_config() monkeypatch.setenv("EMPIRIOLABS_API_KEY", "test-key") - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == EMPIRIOLABS_BASE_URL assert api_key == "test-key" diff --git a/tests/unit/llms/openai_like/test_json_loader.py b/tests/unit/llms/openai_like/test_json_loader.py index 9a97dbcc93d..f6d412fe304 100644 --- a/tests/unit/llms/openai_like/test_json_loader.py +++ b/tests/unit/llms/openai_like/test_json_loader.py @@ -36,7 +36,7 @@ def test_crusoe_get_openai_compatible_provider_info(): # Test with default values (no env vars set) with mock.patch.dict(os.environ, {}, clear=True): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == CRUSOE_API_BASE assert api_key is None @@ -48,7 +48,7 @@ def test_crusoe_get_openai_compatible_provider_info(): "CRUSOE_API_BASE": "https://custom.crusoecloud.com/v1", }, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://custom.crusoecloud.com/v1" assert api_key == "test-key" @@ -60,7 +60,7 @@ def test_crusoe_get_openai_compatible_provider_info(): "CRUSOE_API_BASE": "https://env.crusoecloud.com/v1", }, ): - api_base, api_key = config._get_openai_compatible_provider_info("https://param.crusoecloud.com/v1", "param-key") + api_base, api_key = config.get_openai_compatible_provider_info("https://param.crusoecloud.com/v1", "param-key") assert api_base == "https://param.crusoecloud.com/v1" assert api_key == "param-key" diff --git a/tests/unit/llms/openai_like/test_json_providers.py b/tests/unit/llms/openai_like/test_json_providers.py index a56108ca9ac..9dbdaa28a76 100644 --- a/tests/unit/llms/openai_like/test_json_providers.py +++ b/tests/unit/llms/openai_like/test_json_providers.py @@ -46,13 +46,11 @@ class TestJSONProviderLoader: config = config_class() # Test API info resolution - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://api.publicai.co/v1" # Test with custom base - api_base, api_key = config._get_openai_compatible_provider_info( - "https://custom.api.com", "test-key" - ) + api_base, api_key = config.get_openai_compatible_provider_info("https://custom.api.com", "test-key") assert api_base == "https://custom.api.com" assert api_key == "test-key" @@ -213,12 +211,10 @@ class TestPinstripes: config_class = create_config_class(provider) config = config_class() - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://pinstripes.io/v1" - api_base, api_key = config._get_openai_compatible_provider_info( - "https://custom.pinstripes.io/v1", "test-key" - ) + api_base, api_key = config.get_openai_compatible_provider_info("https://custom.pinstripes.io/v1", "test-key") assert api_base == "https://custom.pinstripes.io/v1" assert api_key == "test-key" @@ -277,12 +273,10 @@ class TestDarkbloom: config_class = create_config_class(provider) config = config_class() - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://api.darkbloom.dev/v1" - api_base, api_key = config._get_openai_compatible_provider_info( - "https://custom.darkbloom.dev/v1", "test-key" - ) + api_base, api_key = config.get_openai_compatible_provider_info("https://custom.darkbloom.dev/v1", "test-key") assert api_base == "https://custom.darkbloom.dev/v1" assert api_key == "test-key" diff --git a/tests/unit/llms/parasail/test_parasail.py b/tests/unit/llms/parasail/test_parasail.py index 8fb9b22b5f6..d37df2e6b74 100644 --- a/tests/unit/llms/parasail/test_parasail.py +++ b/tests/unit/llms/parasail/test_parasail.py @@ -41,7 +41,7 @@ def test_parasail_dynamic_config_env_vars(): "PARASAIL_API_BASE": PARASAIL_RESPONSES_GATEWAY, }, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == PARASAIL_RESPONSES_GATEWAY assert api_key == "test-key" diff --git a/tests/unit/llms/perplexity/chat/test_transformation.py b/tests/unit/llms/perplexity/chat/test_transformation.py index 54ffa13d404..6a70fd1232d 100644 --- a/tests/unit/llms/perplexity/chat/test_transformation.py +++ b/tests/unit/llms/perplexity/chat/test_transformation.py @@ -145,7 +145,7 @@ class TestPerplexityReasoning: from litellm.llms.perplexity.chat.transformation import PerplexityChatConfig config = PerplexityChatConfig() - api_base, _ = config._get_openai_compatible_provider_info(api_base=None, api_key="test-key") + api_base, _ = config.get_openai_compatible_provider_info(api_base=None, api_key="test-key") assert api_base == expected_api_base diff --git a/tests/unit/llms/ragflow/chat/test_ragflow_chat_transformation.py b/tests/unit/llms/ragflow/chat/test_ragflow_chat_transformation.py index 437f53fea1a..1643b76309d 100644 --- a/tests/unit/llms/ragflow/chat/test_ragflow_chat_transformation.py +++ b/tests/unit/llms/ragflow/chat/test_ragflow_chat_transformation.py @@ -382,13 +382,11 @@ class TestRAGFlowChatTransformation: api_base = "http://localhost:9380" api_key = "test-key" - result_api_base, result_api_key, result_provider = ( - config._get_openai_compatible_provider_info( - model=model, - api_base=api_base, - api_key=api_key, - custom_llm_provider="ragflow", - ) + result_api_base, result_api_key, result_provider = config.get_openai_compatible_provider_info( + model=model, + api_base=api_base, + api_key=api_key, + custom_llm_provider="ragflow", ) assert result_api_base == api_base @@ -405,13 +403,11 @@ class TestRAGFlowChatTransformation: model = "ragflow/agent/my-agent-id/gpt-4o-mini" - result_api_base, result_api_key, result_provider = ( - config._get_openai_compatible_provider_info( - model=model, - api_base=None, - api_key=None, - custom_llm_provider="ragflow", - ) + result_api_base, result_api_key, result_provider = config.get_openai_compatible_provider_info( + model=model, + api_base=None, + api_key=None, + custom_llm_provider="ragflow", ) assert result_api_base == "http://env-base:9380" diff --git a/tests/unit/llms/sambanova/test_chat.py b/tests/unit/llms/sambanova/test_chat.py index 1c5d8c5d15d..bf77595757b 100644 --- a/tests/unit/llms/sambanova/test_chat.py +++ b/tests/unit/llms/sambanova/test_chat.py @@ -32,7 +32,7 @@ class TestSambanovaContentListHandling: } ] - transformed_messages = config._transform_messages( + transformed_messages = config.transform_messages( messages=messages, model="sambanova/gpt-oss-120b", is_async=False ) @@ -57,7 +57,7 @@ class TestSambanovaContentListHandling: } ] - transformed_messages = config._transform_messages( + transformed_messages = config.transform_messages( messages=messages, model="sambanova/gpt-oss-120b", is_async=False ) @@ -71,7 +71,7 @@ class TestSambanovaContentListHandling: messages = [{"role": "user", "content": "Hello, how are you?"}] - transformed_messages = config._transform_messages( + transformed_messages = config.transform_messages( messages=messages, model="sambanova/gpt-oss-120b", is_async=False ) @@ -99,7 +99,7 @@ class TestSambanovaContentListHandling: }, ] - transformed_messages = config._transform_messages( + transformed_messages = config.transform_messages( messages=messages, model="sambanova/gpt-oss-120b", is_async=False ) diff --git a/tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py b/tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py index eadf870cb61..c99b28adc49 100644 --- a/tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py +++ b/tests/unit/llms/soniox/audio_transcription/test_soniox_audio_transcription_transformation.py @@ -239,7 +239,7 @@ class TestTransformAudioTranscriptionResponse: ], } } - resp = cfg._build_response_from_payload(payload) + resp = cfg.build_response_from_payload(payload) assert "Speaker 1:" in resp.text assert "Speaker 2:" in resp.text @@ -253,7 +253,7 @@ class TestTransformAudioTranscriptionResponse: ] } } - resp = cfg._build_response_from_payload(payload) + resp = cfg.build_response_from_payload(payload) assert resp["language"] == "en" def test_should_populate_provided_model_response(self): @@ -262,7 +262,7 @@ class TestTransformAudioTranscriptionResponse: model_response._hidden_params = {"pre": "existing"} payload = {"text": "populated"} - resp = cfg._build_response_from_payload(payload, model_response=model_response) + resp = cfg.build_response_from_payload(payload, model_response=model_response) assert resp is model_response assert resp.text == "populated" assert resp._hidden_params["pre"] == "existing" @@ -274,7 +274,7 @@ class TestTransformAudioTranscriptionResponse: "transcription": {"id": "tx_1"}, "transcript": {"text": "hi", "tokens": []}, } - resp = cfg._build_response_from_payload(payload) + resp = cfg.build_response_from_payload(payload) raw = resp._hidden_params["soniox_raw"] assert raw["transcription"]["id"] == "tx_1" assert raw["transcript"]["text"] == "hi" @@ -295,12 +295,12 @@ class TestTransformAudioTranscriptionResponse: ], } } - resp = cfg._build_response_from_payload(payload) + resp = cfg.build_response_from_payload(payload) assert resp.text == "hello world" def test_should_return_empty_text_for_empty_payload(self): cfg = SonioxAudioTranscriptionConfig() - resp = cfg._build_response_from_payload({}) + resp = cfg.build_response_from_payload({}) assert resp.text == "" def test_should_skip_duration_when_audio_duration_ms_is_invalid(self): @@ -309,7 +309,7 @@ class TestTransformAudioTranscriptionResponse: "transcription": {"audio_duration_ms": "not-a-number"}, "transcript": {"text": "hi", "tokens": []}, } - resp = cfg._build_response_from_payload(payload) + resp = cfg.build_response_from_payload(payload) assert "duration" not in resp.model_dump() @@ -597,7 +597,7 @@ class TestBuildResponseWithResponseFormat: ] } } - resp = cfg._build_response_from_payload(payload, response_format="srt") + resp = cfg.build_response_from_payload(payload, response_format="srt") assert "00:00:00,000 --> " in resp.text assert "Hello world." in resp.text @@ -611,7 +611,7 @@ class TestBuildResponseWithResponseFormat: ] } } - resp = cfg._build_response_from_payload(payload, response_format="vtt") + resp = cfg.build_response_from_payload(payload, response_format="vtt") assert resp.text.startswith("WEBVTT\n") assert "Hello world." in resp.text @@ -626,7 +626,7 @@ class TestBuildResponseWithResponseFormat: ], } } - resp = cfg._build_response_from_payload(payload, response_format="verbose_json") + resp = cfg.build_response_from_payload(payload, response_format="verbose_json") # text should be plain (not SRT/VTT) assert resp.text == "Hello world." # words should be populated @@ -650,7 +650,7 @@ class TestBuildResponseWithResponseFormat: ], } } - resp = cfg._build_response_from_payload(payload, response_format=None) + resp = cfg.build_response_from_payload(payload, response_format=None) assert resp.text == "Hello world." def test_should_fallback_to_plain_text_for_srt_with_no_timestamps(self): @@ -663,7 +663,7 @@ class TestBuildResponseWithResponseFormat: } # SRT requested but tokens have no start_ms/end_ms -> empty SRT # falls back gracefully since group_subtitle_tokens_into_cues skips them - resp = cfg._build_response_from_payload(payload, response_format="srt") + resp = cfg.build_response_from_payload(payload, response_format="srt") # With no timestamp data, SRT rendering produces empty string, # but we still get output because the code checks `tokens` truthiness # before choosing SRT path. Actually the tokens list is truthy but diff --git a/tests/unit/llms/test_file_content_block.py b/tests/unit/llms/test_file_content_block.py index 48e1d92b6cc..4ab4ef8b857 100644 --- a/tests/unit/llms/test_file_content_block.py +++ b/tests/unit/llms/test_file_content_block.py @@ -30,7 +30,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( from litellm.llms.gemini.chat.transformation import GoogleAIStudioGeminiConfig from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) from litellm.types.llms.openai import ( AllMessageValues, @@ -102,10 +102,10 @@ def _explicit_null_file_in_content() -> List[AllMessageValues]: def test_gemini_convert_messages_malformed_file_raises_bad_request(): - """_gemini_convert_messages_with_history should raise BadRequestError (not KeyError) + """gemini_convert_messages_with_history should raise BadRequestError (not KeyError) when a content block has type='file' but no 'file' sub-field.""" with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"): - _gemini_convert_messages_with_history( + gemini_convert_messages_with_history( messages=_malformed(), model="gemini-2.0-flash", ) @@ -114,7 +114,7 @@ def test_gemini_convert_messages_malformed_file_raises_bad_request(): def test_gemini_convert_messages_explicit_null_file_field_raises_bad_request(): """Explicit JSON null for `file` must be rejected like a missing `file` key.""" with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"): - _gemini_convert_messages_with_history( + gemini_convert_messages_with_history( messages=_explicit_null_file_in_content(), model="gemini-2.0-flash", ) @@ -130,16 +130,14 @@ def test_google_ai_studio_transform_messages_malformed_file_raises_bad_request() when a content block has type='file' but no 'file' sub-field.""" config = GoogleAIStudioGeminiConfig() with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"): - config._transform_messages(messages=_malformed(), model="gemini-2.0-flash") + config.transform_messages(messages=_malformed(), model="gemini-2.0-flash") def test_google_ai_studio_transform_messages_explicit_null_file_field_raises_bad_request(): """Explicit JSON null for `file` must be rejected like a missing `file` key.""" config = GoogleAIStudioGeminiConfig() with pytest.raises(litellm.BadRequestError, match="missing the required 'file' field"): - config._transform_messages( - messages=_explicit_null_file_in_content(), model="gemini-2.0-flash" - ) + config.transform_messages(messages=_explicit_null_file_in_content(), model="gemini-2.0-flash") def test_google_ai_studio_transform_messages_http_file_id_converts_to_base64(monkeypatch): @@ -176,7 +174,7 @@ def test_google_ai_studio_transform_messages_http_file_id_converts_to_base64(mon ], ) config = GoogleAIStudioGeminiConfig() - config._transform_messages(messages=messages, model="gemini-2.0-flash") + config.transform_messages(messages=messages, model="gemini-2.0-flash") content = messages[0].get("content") assert isinstance(content, list) file_block = next(c for c in content if isinstance(c, dict) and c.get("type") == "file") @@ -219,7 +217,7 @@ def test_google_ai_studio_transform_messages_http_file_id_convert_failure_leaves ], ) config = GoogleAIStudioGeminiConfig() - config._transform_messages(messages=messages, model="gemini-2.0-flash") + config.transform_messages(messages=messages, model="gemini-2.0-flash") content = messages[0].get("content") assert isinstance(content, list) file_block = next(c for c in content if isinstance(c, dict) and c.get("type") == "file") diff --git a/tests/unit/llms/test_lifecycle_fix.py b/tests/unit/llms/test_lifecycle_fix.py index 40362611245..75702b8d666 100644 --- a/tests/unit/llms/test_lifecycle_fix.py +++ b/tests/unit/llms/test_lifecycle_fix.py @@ -17,7 +17,7 @@ async def test_httpx_client_not_closed_by_handler_gc(): After the fix: returns a standalone httpx.AsyncClient, no handler involved. """ # Get the client the same way AsyncOpenAI would - client = BaseOpenAILLM._get_async_http_client() + client = BaseOpenAILLM.get_async_http_client() assert isinstance(client, httpx.AsyncClient) # Simulate what the old code did: create an AsyncHTTPHandler and GC it diff --git a/tests/unit/llms/vercel_ai_gateway/chat/test_vercel_ai_gateway_transformation.py b/tests/unit/llms/vercel_ai_gateway/chat/test_vercel_ai_gateway_transformation.py index a58942559fd..3842dbff7be 100644 --- a/tests/unit/llms/vercel_ai_gateway/chat/test_vercel_ai_gateway_transformation.py +++ b/tests/unit/llms/vercel_ai_gateway/chat/test_vercel_ai_gateway_transformation.py @@ -77,7 +77,7 @@ def test_vercel_ai_gateway_get_openai_compatible_provider_info(): "VERCEL_AI_GATEWAY_API_KEY": "env_api_key", }, ): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://env.vercel.sh/v1" assert api_key == "env_api_key" diff --git a/tests/unit/llms/vertex_ai/batches/test_handler.py b/tests/unit/llms/vertex_ai/batches/test_handler.py index 24df3214da3..59655d5f80c 100644 --- a/tests/unit/llms/vertex_ai/batches/test_handler.py +++ b/tests/unit/llms/vertex_ai/batches/test_handler.py @@ -16,7 +16,7 @@ We mock only true I/O / auth seams: header is forwarded. * ``_check_custom_proxy`` - returns ``(None, url)``; we let it pass the computed default url straight through so we can assert the request URL. - * the httpx client factories (``_get_httpx_client`` / + * the httpx client factories (``get_httpx_client`` / ``get_async_httpx_client``) and the SSRF wrappers (``safe_get`` / ``async_safe_get``) - the network calls. We assert which seam fired with what URL/headers/body, and that the response is parsed into the litellm @@ -119,7 +119,7 @@ def test_create_batch_sync_posts_and_parses(): client = MagicMock() client.post.return_value = _http_response() - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): out = h.create_batch( _is_async=False, create_batch_data=CREATE_DATA, @@ -154,7 +154,7 @@ def test_create_batch_async_returns_coroutine_and_uses_async_client(): sync_client = MagicMock() with ( - patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.get_httpx_client", return_value=sync_client), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), ): coro = h.create_batch( @@ -185,7 +185,7 @@ def test_create_batch_sync_does_not_resolve_publisher_models(): client.post.return_value = _http_response() with ( - patch(f"{HMOD}._get_httpx_client", return_value=client), + patch(f"{HMOD}.get_httpx_client", return_value=client), patch(f"{HMOD}.safe_get") as safe_get, ): out = h.create_batch( @@ -232,7 +232,7 @@ def test_create_batch_sync_resolves_fine_tuned_endpoint_to_tuned_model(): client.post.return_value = _http_response() with ( - patch(f"{HMOD}._get_httpx_client", return_value=client), + patch(f"{HMOD}.get_httpx_client", return_value=client), patch(f"{HMOD}.safe_get", return_value=_endpoint_get_response()) as safe_get, ): out = h.create_batch( @@ -298,7 +298,7 @@ def test_create_batch_sync_endpoint_resolution_error_raises(): resolve_response.text = "endpoint not found" with ( - patch(f"{HMOD}._get_httpx_client", return_value=client), + patch(f"{HMOD}.get_httpx_client", return_value=client), patch(f"{HMOD}.safe_get", return_value=resolve_response), ): with pytest.raises(VertexAIError) as exc_info: @@ -323,7 +323,7 @@ def test_create_batch_custom_endpoint_raises_400_without_io(): h = _make_handler() client = MagicMock() - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): with pytest.raises(VertexAIError) as exc_info: h.create_batch( _is_async=False, @@ -348,7 +348,7 @@ def test_create_batch_sync_endpoint_without_deployed_model_raises_400(): client = MagicMock() with ( - patch(f"{HMOD}._get_httpx_client", return_value=client), + patch(f"{HMOD}.get_httpx_client", return_value=client), patch(f"{HMOD}.safe_get", return_value=_endpoint_get_response(deployed_models=[])), ): with pytest.raises(VertexAIError) as exc_info: @@ -379,7 +379,7 @@ def test_create_batch_sync_httpstatuserror_propagates(): "boom", request=request, response=err_response ) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): with pytest.raises(httpx.HTTPStatusError): h.create_batch( _is_async=False, @@ -398,7 +398,7 @@ def test_create_batch_input_file_id_without_model_raises_400_before_post(): h = _make_handler() client = MagicMock() - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): with pytest.raises(VertexAIError) as exc_info: h.create_batch( _is_async=False, @@ -426,7 +426,7 @@ def test_retrieve_batch_sync_uses_safe_get_with_batch_id_url(): sync_client = MagicMock() with ( - patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.get_httpx_client", return_value=sync_client), patch(f"{HMOD}.safe_get", return_value=_http_response()) as safe_get, ): out = h.retrieve_batch( @@ -456,7 +456,7 @@ def test_retrieve_batch_async_returns_coroutine_uses_async_safe_get(): async_client = MagicMock() with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), patch( f"{HMOD}.async_safe_get", @@ -485,7 +485,7 @@ def test_retrieve_batch_async_returns_coroutine_uses_async_safe_get(): def test_retrieve_batch_sync_non_200_raises(): h = _make_handler() with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.safe_get", return_value=_http_response(status_code=404)), ): with pytest.raises(VertexAIError, match="Error: 404"): @@ -510,7 +510,7 @@ def test_retrieve_batch_sync_invokes_logging_pre_call(): logging_obj = MagicMock(spec=LiteLLMLogging) with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.safe_get", return_value=_http_response()), ): h.retrieve_batch( @@ -549,7 +549,7 @@ def test_list_batches_sync_passes_pagination_params(): client = MagicMock() client.get.return_value = _http_response(json_body=_list_response()) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): out = h.list_batches( _is_async=False, after="cursor-xyz", @@ -578,7 +578,7 @@ def test_list_batches_sync_omits_unset_pagination_params(): client = MagicMock() client.get.return_value = _http_response(json_body={"batchPredictionJobs": []}) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): out = h.list_batches( _is_async=False, after=None, @@ -606,7 +606,7 @@ def test_list_batches_async_returns_coroutine(): sync_client = MagicMock() with ( - patch(f"{HMOD}._get_httpx_client", return_value=sync_client), + patch(f"{HMOD}.get_httpx_client", return_value=sync_client), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), ): coro = h.list_batches( @@ -633,7 +633,7 @@ def test_list_batches_sync_non_200_raises(): client = MagicMock() client.get.return_value = _http_response(status_code=500) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): with pytest.raises(VertexAIError, match="Error: 500"): h.list_batches( _is_async=False, @@ -661,7 +661,7 @@ def test_cancel_batch_sync_posts_cancel_then_retrieves(): json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") ) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): out = h.cancel_batch( _is_async=False, batch_id=BATCH_ID, @@ -697,7 +697,7 @@ def test_cancel_batch_async_returns_coroutine_posts_then_retrieves(): ) with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), ): coro = h.cancel_batch( @@ -726,7 +726,7 @@ def test_cancel_batch_sync_retrieve_non_200_raises(): client.post.return_value = _http_response(json_body={}) client.get.return_value = _http_response(status_code=404) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): with pytest.raises(VertexAIError, match="Error: 404"): h.cancel_batch( _is_async=False, @@ -755,7 +755,7 @@ def test_cancel_batch_sync_proxy_url_without_cancel_suffix_uses_rsplit_branch(): json_body=_vertex_job_response(state="JOB_STATE_CANCELLED") ) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): out = h.cancel_batch( _is_async=False, batch_id=BATCH_ID, @@ -783,7 +783,7 @@ def test_cancel_batch_sync_httpstatuserror_logged_and_reraised(): "boom", request=request, response=err_response ) - with patch(f"{HMOD}._get_httpx_client", return_value=client): + with patch(f"{HMOD}.get_httpx_client", return_value=client): with pytest.raises(httpx.HTTPStatusError): h.cancel_batch( _is_async=False, @@ -810,7 +810,7 @@ def test_create_batch_async_httpstatuserror_logged_and_reraised(): ) with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), ): coro = h.create_batch( @@ -830,7 +830,7 @@ def test_create_batch_async_httpstatuserror_logged_and_reraised(): def test_async_retrieve_batch_non_200_raises(): h = _make_handler() with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=MagicMock()), patch( f"{HMOD}.async_safe_get", @@ -858,7 +858,7 @@ def test_async_retrieve_batch_invokes_logging_pre_call(): logging_obj = MagicMock(spec=LiteLLMLogging) with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=MagicMock()), patch( f"{HMOD}.async_safe_get", @@ -887,7 +887,7 @@ def test_async_list_batches_non_200_raises(): async_client.get = AsyncMock(return_value=_http_response(status_code=500)) with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), ): coro = h.list_batches( @@ -919,7 +919,7 @@ def test_async_cancel_batch_httpstatuserror_and_retrieve_non_200(): ) async_client.get = AsyncMock() with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client), ): coro = h.cancel_batch( @@ -941,7 +941,7 @@ def test_async_cancel_batch_httpstatuserror_and_retrieve_non_200(): async_client2.post = AsyncMock(return_value=_http_response(json_body={})) async_client2.get = AsyncMock(return_value=_http_response(status_code=404)) with ( - patch(f"{HMOD}._get_httpx_client", return_value=MagicMock()), + patch(f"{HMOD}.get_httpx_client", return_value=MagicMock()), patch(f"{HMOD}.get_async_httpx_client", return_value=async_client2), ): coro = h.cancel_batch( diff --git a/tests/unit/llms/vertex_ai/batches/test_transformation.py b/tests/unit/llms/vertex_ai/batches/test_transformation.py index 5f969620a80..cf5e11b35b1 100644 --- a/tests/unit/llms/vertex_ai/batches/test_transformation.py +++ b/tests/unit/llms/vertex_ai/batches/test_transformation.py @@ -24,7 +24,7 @@ from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402 ) from litellm.llms.vertex_ai.common_utils import ( # noqa: E402 VertexAIError, - _convert_vertex_datetime_to_openai_datetime, + convert_vertex_datetime_to_openai_datetime, ) from litellm.types.utils import LiteLLMBatch # noqa: E402 @@ -145,7 +145,7 @@ def test_transform_vertex_response_full_mapping(): # created_at is parsed via the shared helper (uses local tz); assert the # transform forwards createTime through that helper rather than a hardcoded # epoch that would be tz-dependent - assert batch.created_at == _convert_vertex_datetime_to_openai_datetime("2024-12-04T21:53:12.120184Z") + assert batch.created_at == convert_vertex_datetime_to_openai_datetime("2024-12-04T21:53:12.120184Z") assert batch.endpoint == "" assert batch.object == "batch" assert batch.input_file_id == "gs://bucket/in.jsonl" @@ -202,24 +202,24 @@ def test_status_mapping_unknown_state_raises_keyerror(): # =========================================================================== # -# _get_batch_id_from_vertex_ai_batch_response +# get_batch_id_from_vertex_ai_batch_response # =========================================================================== # def test_get_batch_id_splits_path(): assert ( - T._get_batch_id_from_vertex_ai_batch_response({"name": "projects/p/locations/l/batchPredictionJobs/999"}) + T.get_batch_id_from_vertex_ai_batch_response({"name": "projects/p/locations/l/batchPredictionJobs/999"}) == "999" ) def test_get_batch_id_no_slash_returns_name(): - assert T._get_batch_id_from_vertex_ai_batch_response({"name": "abc"}) == "abc" + assert T.get_batch_id_from_vertex_ai_batch_response({"name": "abc"}) == "abc" def test_get_batch_id_empty_name_returns_empty(): - assert T._get_batch_id_from_vertex_ai_batch_response({"name": ""}) == "" - assert T._get_batch_id_from_vertex_ai_batch_response({}) == "" + assert T.get_batch_id_from_vertex_ai_batch_response({"name": ""}) == "" + assert T.get_batch_id_from_vertex_ai_batch_response({}) == "" # =========================================================================== # diff --git a/tests/unit/llms/vertex_ai/gemini/test_context_circulation.py b/tests/unit/llms/vertex_ai/gemini/test_context_circulation.py index 1422531edbf..39c40eeae67 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_context_circulation.py +++ b/tests/unit/llms/vertex_ai/gemini/test_context_circulation.py @@ -14,12 +14,12 @@ from unittest.mock import MagicMock import pytest +from litellm.llms.vertex_ai.gemini.transformation import ( + gemini_convert_messages_with_history, +) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) -from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, -) from litellm.types.llms.vertex_ai import HttpxPartType # --- Response extraction tests --- @@ -227,7 +227,7 @@ class TestReInjectServerSideToolInvocations: {"role": "user", "content": "Thanks!"}, ] - contents = _gemini_convert_messages_with_history(messages) + contents = gemini_convert_messages_with_history(messages) # Find the model turn model_turn = [c for c in contents if c["role"] == "model"] @@ -276,7 +276,7 @@ class TestReInjectServerSideToolInvocations: {"role": "user", "content": "Thanks!"}, ] - contents = _gemini_convert_messages_with_history(messages) + contents = gemini_convert_messages_with_history(messages) model_turn = [c for c in contents if c["role"] == "model"] assert len(model_turn) == 1 @@ -312,7 +312,7 @@ class TestReInjectServerSideToolInvocations: {"role": "user", "content": "Thanks!"}, ] - contents = _gemini_convert_messages_with_history(messages) + contents = gemini_convert_messages_with_history(messages) model_turn = [c for c in contents if c["role"] == "model"] assert len(model_turn) == 1 @@ -331,7 +331,7 @@ class TestReInjectServerSideToolInvocations: {"role": "user", "content": "Bye"}, ] - contents = _gemini_convert_messages_with_history(messages) + contents = gemini_convert_messages_with_history(messages) model_turn = [c for c in contents if c["role"] == "model"] assert len(model_turn) == 1 diff --git a/tests/unit/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py b/tests/unit/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py index 10fc68ecaad..f649e011dfe 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py +++ b/tests/unit/llms/vertex_ai/gemini/test_gemini_image_url_missing_field.py @@ -1,9 +1,10 @@ -import pytest from typing import List, cast +import pytest + import litellm from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) from litellm.types.llms.openai import AllMessageValues @@ -15,7 +16,7 @@ def test_missing_image_url_field_raises_bad_request_error(): [{"role": "user", "content": [{"type": "image_url"}]}], ) with pytest.raises(litellm.BadRequestError) as exc_info: - _gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") + gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") assert "'image_url' field is missing" in str(exc_info.value) @@ -26,7 +27,7 @@ def test_missing_url_inside_image_url_dict_raises_bad_request_error(): [{"role": "user", "content": [{"type": "image_url", "image_url": {"detail": "high"}}]}], ) with pytest.raises(litellm.BadRequestError) as exc_info: - _gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") + gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") assert "'url' field is missing inside" in str(exc_info.value) @@ -37,7 +38,7 @@ def test_explicit_null_image_url_raises_bad_request_error(): [{"role": "user", "content": [{"type": "image_url", "image_url": None}]}], ) with pytest.raises(litellm.BadRequestError) as exc_info: - _gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") + gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") assert "'image_url' field is missing" in str(exc_info.value) @@ -48,5 +49,5 @@ def test_empty_dict_image_url_raises_bad_request_error(): [{"role": "user", "content": [{"type": "image_url", "image_url": {}}]}], ) with pytest.raises(litellm.BadRequestError) as exc_info: - _gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") + gemini_convert_messages_with_history(messages, model="gemini-1.5-pro") assert "'url' field is missing inside" in str(exc_info.value) diff --git a/tests/unit/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py b/tests/unit/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py index bbd12e25f43..a292b712a49 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py +++ b/tests/unit/llms/vertex_ai/gemini/test_tool_call_followed_by_text_assistant.py @@ -16,7 +16,7 @@ This happens for any OpenAI-style history with that shape, independent of provid import pytest from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) @@ -44,7 +44,7 @@ def test_tool_result_matches_tool_call_with_text_assistant_in_between(): messages = _messages_with_text_assistant_between_tool_call_and_result() # Should not raise "Missing corresponding tool call for tool response message". - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # The function response must be present and carry the correct tool name. function_responses = [ diff --git a/tests/unit/llms/vertex_ai/gemini/test_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_transformation.py index f135acd094f..2ea02b776ab 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini/test_transformation.py @@ -42,7 +42,7 @@ async def test__transform_request_body_labels(): "cached_content": None, } - rb: RequestBody = transformation._transform_request_body(**transform_request_params) + rb: RequestBody = transformation.transform_request_body(**transform_request_params) # Check URL assert rb["contents"] == [ @@ -83,7 +83,7 @@ async def test__transform_request_body_metadata(): "cached_content": None, } - rb: RequestBody = transformation._transform_request_body(**transform_request_params) + rb: RequestBody = transformation.transform_request_body(**transform_request_params) # Check URL assert rb["contents"] == [ @@ -126,7 +126,7 @@ async def test__transform_request_body_labels_and_metadata(): "cached_content": None, } - rb: RequestBody = transformation._transform_request_body(**transform_request_params) + rb: RequestBody = transformation.transform_request_body(**transform_request_params) # Check URL assert rb["contents"] == [ @@ -171,7 +171,7 @@ async def test__transform_request_body_image_config(): "cached_content": None, } - rb: RequestBody = transformation._transform_request_body(**transform_request_params) + rb: RequestBody = transformation.transform_request_body(**transform_request_params) assert "generationConfig" in rb assert "imageConfig" in rb["generationConfig"] @@ -207,7 +207,7 @@ async def test__transform_request_body_image_config_snake_case(): "cached_content": None, } - rb: RequestBody = transformation._transform_request_body(**transform_request_params) + rb: RequestBody = transformation.transform_request_body(**transform_request_params) assert "generationConfig" in rb assert "image_config" in rb["generationConfig"] @@ -240,7 +240,7 @@ async def test__transform_request_body_image_config_with_image_size(): "cached_content": None, } - rb: RequestBody = transformation._transform_request_body(**transform_request_params) + rb: RequestBody = transformation.transform_request_body(**transform_request_params) assert "generationConfig" in rb assert "imageConfig" in rb["generationConfig"] @@ -270,7 +270,7 @@ def test__transform_request_body_google_maps_json_schema_uses_response_format(): "cached_content": None, } - rb: RequestBody = transformation._transform_request_body(**transform_request_params) + rb: RequestBody = transformation.transform_request_body(**transform_request_params) gen = rb["generationConfig"] assert "responseFormat" in gen @@ -290,7 +290,7 @@ def test_map_function_google_search_snake_case(): # Test snake_case google_search tools = [{"google_search": {}}] - result = config._map_function(tools, optional_params) + result = config.map_function(tools, optional_params) assert len(result) == 1 assert "googleSearch" in result[0] @@ -306,7 +306,7 @@ def test_map_function_google_search_camel_case(): # Test camelCase googleSearch tools = [{"googleSearch": {}}] - result = config._map_function(tools, optional_params) + result = config.map_function(tools, optional_params) assert len(result) == 1 assert "googleSearch" in result[0] @@ -320,14 +320,8 @@ def test_map_function_google_search_retrieval_snake_case(): config = VertexGeminiConfig() optional_params = {} - tools = [ - { - "google_search_retrieval": { - "dynamic_retrieval_config": {"mode": "MODE_DYNAMIC"} - } - } - ] - result = config._map_function(tools, optional_params) + tools = [{"google_search_retrieval": {"dynamic_retrieval_config": {"mode": "MODE_DYNAMIC"}}}] + result = config.map_function(tools, optional_params) assert len(result) == 1 assert "googleSearchRetrieval" in result[0] @@ -341,7 +335,7 @@ def test_map_function_enterprise_web_search_snake_case(): optional_params = {} tools = [{"enterprise_web_search": {}}] - result = config._map_function(tools, optional_params) + result = config.map_function(tools, optional_params) assert len(result) == 1 assert "enterpriseWebSearch" in result[0] diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 114655e6277..047af5d0ab2 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -13,11 +13,11 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( convert_to_gemini_tool_call_result, ) from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - _transform_request_body, - check_if_part_exists_in_parts, - _get_highest_media_resolution, _extract_max_media_resolution_from_messages, + _get_highest_media_resolution, + check_if_part_exists_in_parts, + gemini_convert_messages_with_history, + transform_request_body, ) from litellm.types.llms.vertex_ai import BlobType, ContentType, PartType from litellm.types.utils import Message @@ -121,7 +121,7 @@ def test_cached_content_respects_modify_params_for_cache_incompatible_fields(): try: # With modify_params=False (default), keep fields even with cachedContent. litellm.modify_params = False - result = _transform_request_body( + result = transform_request_body( messages=list(messages), model="gemini-2.5-pro", optional_params=dict(optional_params), @@ -137,7 +137,7 @@ def test_cached_content_respects_modify_params_for_cache_incompatible_fields(): # With modify_params=True, drop cache-incompatible fields. litellm.modify_params = True - result_modify_true = _transform_request_body( + result_modify_true = transform_request_body( messages=list(messages), model="gemini-2.5-pro", optional_params=dict(optional_params), @@ -152,7 +152,7 @@ def test_cached_content_respects_modify_params_for_cache_incompatible_fields(): assert "contents" in result_modify_true # Without cache, fields are always included. - result_no_cache = _transform_request_body( + result_no_cache = transform_request_body( messages=list(messages), model="gemini-2.5-pro", optional_params=dict(optional_params), @@ -174,7 +174,7 @@ def test_google_genai_excludes_labels(): optional_params = {"labels": {"project": "test", "team": "ai"}} litellm_params = {} - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params, @@ -194,7 +194,7 @@ def test_vertex_ai_includes_labels(): optional_params = {"labels": {"project": "test", "team": "ai"}} litellm_params = {} - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params, @@ -214,7 +214,7 @@ def test_service_tier_forwarded_to_vertex_ai(): optional_params = {"service_tier": "flex"} litellm_params = {} - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params, @@ -244,7 +244,7 @@ def test_extra_body_cache_not_forwarded_to_vertex_ai(): } litellm_params = {} - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params, @@ -282,7 +282,7 @@ def test_extra_body_tags_not_forwarded_to_vertex_ai(): } litellm_params = {} - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params, @@ -309,7 +309,7 @@ def test_extra_body_google_maps_rewrites_json_response_format(): }, } - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params, @@ -347,7 +347,7 @@ def test_extra_body_generation_config_cannot_restore_google_maps_json_mime_type( }, } - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params, @@ -374,14 +374,10 @@ def test_metadata_to_labels_vertex_only(): """Test that metadata->labels conversion only happens for Vertex AI""" messages = [{"role": "user", "content": "test"}] optional_params = {} - litellm_params = { - "metadata": { - "requester_metadata": {"user": "john_doe", "project": "test-project"} - } - } + litellm_params = {"metadata": {"requester_metadata": {"user": "john_doe", "project": "test-project"}}} # Google GenAI/AI Studio should not include labels from metadata - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params.copy(), @@ -392,7 +388,7 @@ def test_metadata_to_labels_vertex_only(): assert "labels" not in result # Vertex AI should include labels from metadata - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-pro", optional_params=optional_params.copy(), @@ -409,7 +405,7 @@ def test_empty_content_handling(): # Test with empty content in user message messages = [{"content": "", "role": "user"}] - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # Verify that the content was properly transformed assert len(contents) == 1 @@ -1032,7 +1028,7 @@ def test_real_signature_forwarded_to_gemini_2_5_without_placeholder_siblings(): def test_parallel_tool_call_history_replayed_through_full_message_conversion(): """End to end through the message-history converter, the path a real /chat/completions replay takes.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -1040,15 +1036,11 @@ def test_parallel_tool_call_history_replayed_through_full_message_conversion(): { "role": "assistant", "content": None, - "tool_calls": _parallel_tool_calls_signed_via_id( - REAL_THOUGHT_SIGNATURE, None, None - ), + "tool_calls": _parallel_tool_calls_signed_via_id(REAL_THOUGHT_SIGNATURE, None, None), }, ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") model_parts = contents[1]["parts"] assert len(model_parts) == 3 @@ -1070,7 +1062,7 @@ def test_natively_signed_parallel_turn_never_carries_a_placeholder(model): import json from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -1078,13 +1070,11 @@ def test_natively_signed_parallel_turn_never_carries_a_placeholder(model): { "role": "assistant", "content": None, - "tool_calls": _parallel_tool_calls_signed_via_id( - REAL_THOUGHT_SIGNATURE, None, None - ), + "tool_calls": _parallel_tool_calls_signed_via_id(REAL_THOUGHT_SIGNATURE, None, None), }, ] - contents = _gemini_convert_messages_with_history(messages=messages, model=model) + contents = gemini_convert_messages_with_history(messages=messages, model=model) model_parts = contents[1]["parts"] assert len(model_parts) == 3 @@ -1138,7 +1128,7 @@ def test_signed_text_part_survives_alongside_unsigned_parallel_tool_calls(): """Text-part and function-call signatures are collected by separate code paths, so scoping the placeholder must not disturb a real signature that arrived on the text part.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) msg = { @@ -1148,9 +1138,7 @@ def test_signed_text_part_survives_alongside_unsigned_parallel_tool_calls(): "tool_calls": _parallel_tool_calls(None, None, None), } - parts = _gemini_convert_messages_with_history( - messages=[msg], model="gemini-3-pro-preview" - )[0]["parts"] + parts = gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro-preview")[0]["parts"] assert parts[0]["text"] == "Checking all three cities." assert parts[0]["thoughtSignature"] == "real_25_signature" @@ -1305,7 +1293,7 @@ class TestMediaResolution: } ] - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-flash", optional_params={}, @@ -1336,7 +1324,7 @@ class TestMediaResolution: } ] - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-flash", optional_params={}, @@ -1367,7 +1355,7 @@ class TestMediaResolution: } ] - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-3-pro-preview", optional_params={}, @@ -1396,7 +1384,7 @@ class TestMediaResolution: } ] - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-2.5-flash", optional_params={}, @@ -1472,7 +1460,7 @@ class TestMediaResolution: } ] - result = _transform_request_body( + result = transform_request_body( messages=messages, model="gemini-1.5-pro", optional_params={}, @@ -1517,9 +1505,7 @@ class TestVideoMetadataAllGeminiModels: def test_video_metadata_fps_gemini_2_5_flash(self): """Gemini 2.5 Flash: fps in video_metadata should be forwarded (Issue #25474)""" messages = self._make_video_messages({"fps": 5}) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-flash" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-2.5-flash") file_part = self._get_file_part(contents) assert "video_metadata" in file_part assert file_part["video_metadata"]["fps"] == 5 @@ -1527,21 +1513,15 @@ class TestVideoMetadataAllGeminiModels: def test_video_metadata_fps_gemini_2_5_pro(self): """Gemini 2.5 Pro: fps in video_metadata should be forwarded (Issue #25474)""" messages = self._make_video_messages({"fps": 10}) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-pro" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-2.5-pro") file_part = self._get_file_part(contents) assert "video_metadata" in file_part assert file_part["video_metadata"]["fps"] == 10 def test_video_metadata_offsets_gemini_2_5_flash(self): """Gemini 2.5 Flash: start_offset/end_offset converted to camelCase (Issue #25474)""" - messages = self._make_video_messages( - {"start_offset": "5s", "end_offset": "30s"} - ) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-flash" - ) + messages = self._make_video_messages({"start_offset": "5s", "end_offset": "30s"}) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-2.5-flash") file_part = self._get_file_part(contents) assert "video_metadata" in file_part vm = file_part["video_metadata"] @@ -1550,12 +1530,8 @@ class TestVideoMetadataAllGeminiModels: def test_video_metadata_all_fields_gemini_2_5_flash(self): """Gemini 2.5 Flash: all video_metadata fields forwarded correctly (Issue #25474)""" - messages = self._make_video_messages( - {"fps": 5, "start_offset": "10s", "end_offset": "60s"} - ) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-flash" - ) + messages = self._make_video_messages({"fps": 5, "start_offset": "10s", "end_offset": "60s"}) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-2.5-flash") file_part = self._get_file_part(contents) assert "video_metadata" in file_part vm = file_part["video_metadata"] @@ -1566,9 +1542,7 @@ class TestVideoMetadataAllGeminiModels: def test_video_metadata_gemini_1_5_pro(self): """Gemini 1.5 Pro: video_metadata should also be forwarded (Issue #25474)""" messages = self._make_video_messages({"fps": 2}) - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-1.5-pro" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-1.5-pro") file_part = self._get_file_part(contents) assert "video_metadata" in file_part assert file_part["video_metadata"]["fps"] == 2 @@ -1663,7 +1637,7 @@ def test_gemini_history_nests_multimodal_tool_response_parts(): }, ] - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) tool_response_parts = contents[-1]["parts"] assert len(tool_response_parts) == 1 @@ -2122,7 +2096,7 @@ def test_assistant_message_with_images_field(): ] # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # Verify structure assert len(contents) == 2, f"Expected 2 content blocks, got {len(contents)}" @@ -2192,19 +2166,15 @@ def test_assistant_message_with_multiple_images(): ] # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # Verify assistant message has 3 parts (1 text + 2 images) assert contents[1]["role"] == "model" - assert ( - len(contents[1]["parts"]) == 3 - ), f"Expected 3 parts (text + 2 images), got {len(contents[1]['parts'])}" + assert len(contents[1]["parts"]) == 3, f"Expected 3 parts (text + 2 images), got {len(contents[1]['parts'])}" # Count inline_data parts inline_data_parts = [part for part in contents[1]["parts"] if "inline_data" in part] - assert ( - len(inline_data_parts) == 2 - ), f"Expected 2 inline_data parts, got {len(inline_data_parts)}" + assert len(inline_data_parts) == 2, f"Expected 2 inline_data parts, got {len(inline_data_parts)}" # Verify first image assert inline_data_parts[0]["inline_data"]["mime_type"] == "image/png" @@ -2241,7 +2211,7 @@ def test_assistant_message_with_images_using_message_object(): messages = [user_message, assistant_message] # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # Verify assistant message has both text and image assert contents[1]["role"] == "model" @@ -2283,7 +2253,7 @@ def test_assistant_message_with_images_in_conversation_history(): ] # Convert messages to Gemini format - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # Verify structure: user -> model (with image) -> user assert len(contents) == 3 @@ -2331,7 +2301,7 @@ def test_function_response_has_user_role(): }, ] - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # Expect: user -> model (functionCall) -> user (functionResponse) assert len(contents) == 3 @@ -2401,7 +2371,7 @@ def test_multi_turn_function_calling_roles(): }, ] - contents = _gemini_convert_messages_with_history(messages=messages) + contents = gemini_convert_messages_with_history(messages=messages) # Every content block must have a valid role for i, content in enumerate(contents): @@ -2415,19 +2385,19 @@ def test_multi_turn_function_calling_roles(): for i, content in enumerate(contents): for part in content["parts"]: if "function_response" in part: - assert ( - content["role"] == "user" - ), f"Content block {i} with function_response has role='{content['role']}', expected 'user'" + assert content["role"] == "user", ( + f"Content block {i} with function_response has role='{content['role']}', expected 'user'" + ) def test_gemini_thought_signature_preservation_real_response(): """Test that thought signatures are preserved on the text part if originally there, without dropping or duplicating (real response case).""" + from litellm.llms.vertex_ai.gemini.transformation import ( + gemini_convert_messages_with_history, + ) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, ) - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) real_candidate = { "content": { @@ -2470,11 +2440,9 @@ def test_gemini_thought_signature_preservation_real_response(): if functions is not None: msg["function_call"] = functions if thought_signatures is not None: - msg["provider_specific_fields"] = { - "thought_signatures": thought_signatures - } + msg["provider_specific_fields"] = {"thought_signatures": thought_signatures} - converted_real = _gemini_convert_messages_with_history( + converted_real = gemini_convert_messages_with_history( messages=[msg], model="gemini-2.5-pro", ) @@ -2494,28 +2462,24 @@ def test_gemini_thought_signature_preservation_real_response(): def test_gemini_thought_signature_deduplication_assumed_response(): """Test that thought signatures are deduplicated and not attached to the text part if already present in the tool call (assumed response case).""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) pr_assumed_msg = { "role": "assistant", "content": "I will list the directory.", - "provider_specific_fields": { - "thought_signatures": ["mock_signature_63k"] - }, + "provider_specific_fields": {"thought_signatures": ["mock_signature_63k"]}, "tool_calls": [ { "id": "call_1", "type": "function", "function": {"name": "list_files", "arguments": "{}"}, - "provider_specific_fields": { - "thought_signature": "mock_signature_63k" - }, + "provider_specific_fields": {"thought_signature": "mock_signature_63k"}, } ], } - converted_pr = _gemini_convert_messages_with_history( + converted_pr = gemini_convert_messages_with_history( messages=[pr_assumed_msg], model="gemini-2.5-pro", ) @@ -2533,18 +2497,16 @@ def test_gemini_thought_signature_deduplication_assumed_response(): def test_gemini_thought_signature_pure_text(): """Test that thought signatures are preserved on the text part for responses with no tool calls.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) msg = { "role": "assistant", "content": "Hello, I am a model.", - "provider_specific_fields": { - "thought_signatures": ["pure_text_signature"] - }, + "provider_specific_fields": {"thought_signatures": ["pure_text_signature"]}, } - converted = _gemini_convert_messages_with_history( + converted = gemini_convert_messages_with_history( messages=[msg], model="gemini-2.5-pro", ) @@ -2560,28 +2522,24 @@ def test_gemini_thought_signature_pure_text(): def test_gemini_thought_signature_pure_tool_call(): """Test that thought signatures are preserved on the tool call for responses with no intermediate text.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) msg = { "role": "assistant", "content": None, - "provider_specific_fields": { - "thought_signatures": ["pure_tool_signature"] - }, + "provider_specific_fields": {"thought_signatures": ["pure_tool_signature"]}, "tool_calls": [ { "id": "call_1", "type": "function", "function": {"name": "list_files", "arguments": "{}"}, - "provider_specific_fields": { - "thought_signature": "pure_tool_signature" - }, + "provider_specific_fields": {"thought_signature": "pure_tool_signature"}, } ], } - converted = _gemini_convert_messages_with_history( + converted = gemini_convert_messages_with_history( messages=[msg], model="gemini-2.5-pro", ) @@ -2597,15 +2555,13 @@ def test_gemini_thought_signature_pure_tool_call(): def test_gemini_distinct_text_and_tool_signatures_are_both_preserved(): """A text-part signature that differs from the tool-call signature must stay on the text part.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) msg = { "role": "assistant", "content": "Some analysis.", - "provider_specific_fields": { - "thought_signatures": ["text_signature", "tool_signature"] - }, + "provider_specific_fields": {"thought_signatures": ["text_signature", "tool_signature"]}, "tool_calls": [ { "id": "call_1", @@ -2616,9 +2572,7 @@ def test_gemini_distinct_text_and_tool_signatures_are_both_preserved(): ], } - parts = _gemini_convert_messages_with_history( - messages=[msg], model="gemini-2.5-pro" - )[0]["parts"] + parts = gemini_convert_messages_with_history(messages=[msg], model="gemini-2.5-pro")[0]["parts"] assert parts[0]["text"] == "Some analysis." assert parts[0]["thoughtSignature"] == "text_signature" @@ -2633,7 +2587,7 @@ def test_gemini_25_text_signature_survives_replay_to_gemini_3(): _get_dummy_thought_signature, ) from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) msg = { @@ -2649,9 +2603,7 @@ def test_gemini_25_text_signature_survives_replay_to_gemini_3(): ], } - parts = _gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[ - 0 - ]["parts"] + parts = gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[0]["parts"] assert parts[0]["text"] == "I will list the directory." assert parts[0]["thoughtSignature"] == "real_25_signature" @@ -2663,7 +2615,7 @@ def test_gemini_function_call_signature_round_trip_no_duplicate(): """End to end: a gemini-3-style response (unsigned text + signed functionCall) parsed and re-serialized sends the signature exactly once, on the function-call part.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( VertexGeminiConfig, @@ -2693,9 +2645,7 @@ def test_gemini_function_call_signature_round_trip_no_duplicate(): "provider_specific_fields": {"thought_signatures": thought_signatures}, } - parts = _gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[ - 0 - ]["parts"] + parts = gemini_convert_messages_with_history(messages=[msg], model="gemini-3-pro")[0]["parts"] signatures = [p["thoughtSignature"] for p in parts if "thoughtSignature" in p] assert signatures == ["signature_from_function_call"] @@ -2706,7 +2656,7 @@ def test_gemini_function_call_signature_round_trip_no_duplicate(): def test_gemini_server_side_tool_signature_not_duplicated_on_text(): """A signature already re-injected on a server-side toolCall part is not attached to the text part again.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) msg = { @@ -2726,9 +2676,7 @@ def test_gemini_server_side_tool_signature_not_duplicated_on_text(): }, } - parts = _gemini_convert_messages_with_history( - messages=[msg], model="gemini-2.5-pro" - )[0]["parts"] + parts = gemini_convert_messages_with_history(messages=[msg], model="gemini-2.5-pro")[0]["parts"] text_part = next(p for p in parts if "text" in p) assert "thoughtSignature" not in text_part @@ -2793,7 +2741,7 @@ def test_thinking_block_signature_is_not_forwarded_to_gemini() -> None: {"role": "user", "content": "And of Spain?"}, ] - parts: Final = _parts_of(_gemini_convert_messages_with_history(messages=messages, model="gemini-3.8-flash")) + parts: Final = _parts_of(gemini_convert_messages_with_history(messages=messages, model="gemini-3.8-flash")) assert all("thoughtSignature" not in part for part in parts) assert [part for part in parts if part.get("text") == thinking] == [{"thought": True, "text": thinking}] @@ -2820,7 +2768,7 @@ def test_anthropic_messages_history_replays_to_gemini_without_claude_signature() chat_messages: Final = LiteLLMAnthropicMessagesAdapter().translate_anthropic_messages_to_openai( messages=anthropic_messages ) - parts: Final = _parts_of(_gemini_convert_messages_with_history(messages=chat_messages, model="gemini-3.8-flash")) + parts: Final = _parts_of(gemini_convert_messages_with_history(messages=chat_messages, model="gemini-3.8-flash")) signatures: Final = [part["thoughtSignature"] for part in parts if "thoughtSignature" in part] assert signatures == [_get_dummy_thought_signature()] diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 069b080ebed..02fe48c0e38 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -23,7 +23,7 @@ from litellm.types.utils import ChoiceLogprobs, Usage from litellm.utils import _invalidate_model_cost_lowercase_map, CustomStreamWrapper from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER from litellm.llms.vertex_ai.gemini.transformation import( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome @@ -60,13 +60,13 @@ def test_get_model_for_vertex_ai_url(): def test_is_model_gemini_spec_model(): # Test case 1: None input - assert VertexGeminiConfig._is_model_gemini_spec_model(None) == False + assert VertexGeminiConfig.is_model_gemini_spec_model(None) == False # Test case 2: Regular model name - assert VertexGeminiConfig._is_model_gemini_spec_model("gemini-pro") == False + assert VertexGeminiConfig.is_model_gemini_spec_model("gemini-pro") == False # Test case 3: Gemini spec model - assert VertexGeminiConfig._is_model_gemini_spec_model("gemini/custom-model") == True + assert VertexGeminiConfig.is_model_gemini_spec_model("gemini/custom-model") == True def test_get_model_name_from_gemini_spec_model(): @@ -500,19 +500,14 @@ def test_vertex_ai_empty_content(): ), ], ) -def test_vertex_ai_candidate_token_count_inclusive( - usage_metadata, inclusive, expected_usage -): +def test_vertex_ai_candidate_token_count_inclusive(usage_metadata, inclusive, expected_usage): """ Test that the candidate token count is inclusive of the thinking token count """ v = VertexGeminiConfig() - assert ( - VertexGeminiConfig.is_candidate_token_count_inclusive(usage_metadata) - is inclusive - ) + assert VertexGeminiConfig.is_candidate_token_count_inclusive(usage_metadata) is inclusive - usage = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + usage = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) assert usage.prompt_tokens == expected_usage.prompt_tokens assert usage.completion_tokens == expected_usage.completion_tokens assert usage.total_tokens == expected_usage.total_tokens @@ -534,7 +529,7 @@ def test_vertex_ai_grounded_usage_surfaces_tool_use_tokens(): toolUsePromptTokenCount=12499, ) - usage = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + usage = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) assert usage.prompt_tokens + usage.completion_tokens == usage.total_tokens assert usage.prompt_tokens_details.tool_use_tokens == 12499 @@ -549,7 +544,7 @@ def test_vertex_ai_non_grounded_usage_omits_tool_use_tokens(): totalTokenCount=20, ) - usage = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + usage = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) assert usage.prompt_tokens == 10 assert not hasattr(usage.prompt_tokens_details, "tool_use_tokens") @@ -630,7 +625,7 @@ def test_vertex_ai_maps_grounding_tool_use_tokens_excluded_from_prompt_tokens(): ), } - usage = v._calculate_usage(completion_response=completion_response) + usage = v.calculate_usage(completion_response=completion_response) assert usage.prompt_tokens == 15 assert usage.completion_tokens == 100 @@ -706,7 +701,7 @@ def test_vertex_ai_search_grounding_tool_use_tokens_excluded_from_prompt_tokens( ), } - usage = v._calculate_usage(completion_response=completion_response) + usage = v.calculate_usage(completion_response=completion_response) assert usage.prompt_tokens == 19 assert usage.completion_tokens == 304 + 122 @@ -739,7 +734,7 @@ def test_vertex_ai_url_context_tool_use_tokens_billed_as_input_tokens(): ), } - usage = v._calculate_usage(completion_response=completion_response) + usage = v.calculate_usage(completion_response=completion_response) assert usage.prompt_tokens == 19 + 142 assert usage.completion_tokens == 304 + 122 @@ -928,21 +923,13 @@ def test_streaming_chunk_with_tool_calls_no_thought_no_reasoning_content(): # Tool calls should still work assert streaming_chunk.choices[0].delta.tool_calls is not None assert len(streaming_chunk.choices[0].delta.tool_calls) == 1 - assert ( - streaming_chunk.choices[0].delta.tool_calls[0].function.name - == "get_current_time" - ) + assert streaming_chunk.choices[0].delta.tool_calls[0].function.name == "get_current_time" def test_check_finish_reason(): finish_reason_mappings = VertexGeminiConfig.get_finish_reason_mapping() for k, v in finish_reason_mappings.items(): - assert ( - VertexGeminiConfig._check_finish_reason( - chat_completion_message=None, finish_reason=k - ) - == v - ) + assert VertexGeminiConfig.check_finish_reason(chat_completion_message=None, finish_reason=k) == v def test_every_documented_gemini_finish_reason_has_an_explicit_mapping(): @@ -960,18 +947,14 @@ def test_finish_reason_unspecified_and_malformed_function_call(): # Test FINISH_REASON_UNSPECIFIED maps to "stop" assert finish_reason_mappings["FINISH_REASON_UNSPECIFIED"] == "stop" assert ( - VertexGeminiConfig._check_finish_reason( - chat_completion_message=None, finish_reason="FINISH_REASON_UNSPECIFIED" - ) + VertexGeminiConfig.check_finish_reason(chat_completion_message=None, finish_reason="FINISH_REASON_UNSPECIFIED") == "stop" ) # Test MALFORMED_FUNCTION_CALL maps to "stop" assert finish_reason_mappings["MALFORMED_FUNCTION_CALL"] == "stop" assert ( - VertexGeminiConfig._check_finish_reason( - chat_completion_message=None, finish_reason="MALFORMED_FUNCTION_CALL" - ) + VertexGeminiConfig.check_finish_reason(chat_completion_message=None, finish_reason="MALFORMED_FUNCTION_CALL") == "stop" ) @@ -1002,7 +985,7 @@ def test_vertex_ai_usage_metadata_response_token_count(): "responseTokensDetails": [{"modality": "TEXT", "tokenCount": 74}], } usage_metadata = UsageMetadata(**usage_metadata) - result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + result = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) print("result", result) assert result.prompt_tokens == 66 assert result.completion_tokens == 74 @@ -1033,7 +1016,7 @@ def test_vertex_ai_usage_metadata_with_image_tokens(): "thoughtsTokenCount": 158, } usage_metadata = UsageMetadata(**usage_metadata) - result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + result = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) print("result", result) # Verify basic token counts @@ -1075,7 +1058,7 @@ def test_vertex_ai_usage_metadata_with_image_tokens_auto_calculated_text(): "thoughtsTokenCount": 158, } usage_metadata = UsageMetadata(**usage_metadata) - result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + result = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) print("result", result) # Verify basic token counts @@ -1119,7 +1102,7 @@ def test_vertex_ai_usage_metadata_with_image_tokens_in_prompt(): "thoughtsTokenCount": 217, } usage_metadata = UsageMetadata(**usage_metadata) - result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + result = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) print("result", result) # Verify basic token counts @@ -1175,7 +1158,7 @@ def test_vertex_ai_usage_metadata_accumulates_duplicate_modalities(): ], } usage_metadata = UsageMetadata(**usage_metadata) - result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + result = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) # prompt details are total - cached per modality assert result.prompt_tokens_details.text_tokens == 16 # 20 - 4 @@ -1219,17 +1202,13 @@ def test_vertex_ai_map_thinking_param_without_budget_tokens_for_gemini_3(): def test_vertex_ai_map_tools(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( - value=[{"code_execution": {}}], optional_params=optional_params - ) + tools = v.map_function(value=[{"code_execution": {}}], optional_params=optional_params) assert len(tools) == 1 assert tools[0]["code_execution"] == {} print(tools) new_optional_params = {} - new_tools = v._map_function( - value=[{"codeExecution": {}}], optional_params=new_optional_params - ) + new_tools = v.map_function(value=[{"codeExecution": {}}], optional_params=new_optional_params) assert len(new_tools) == 1 print("new_tools", new_tools) assert new_tools[0]["code_execution"] == {} @@ -1269,13 +1248,13 @@ def test_vertex_ai_map_tool_with_anyof(): }, } ] - tools = v._map_function(value=value, optional_params=optional_params) + tools = v.map_function(value=value, optional_params=optional_params) - assert tools[0]["function_declarations"][0]["parameters"]["properties"][ - "base_branch" - ] == { + assert tools[0]["function_declarations"][0]["parameters"]["properties"]["base_branch"] == { "anyOf": [{"type": "string", "nullable": True, "title": "Base Branch"}] - }, f"Expected only anyOf field and its contents to be kept, but got {tools[0]['function_declarations'][0]['parameters']['properties']['base_branch']}" + }, ( + f"Expected only anyOf field and its contents to be kept, but got {tools[0]['function_declarations'][0]['parameters']['properties']['base_branch']}" + ) new_optional_params = {} new_value = [ @@ -1300,13 +1279,13 @@ def test_vertex_ai_map_tool_with_anyof(): }, } ] - new_tools = v._map_function(value=new_value, optional_params=new_optional_params) + new_tools = v.map_function(value=new_value, optional_params=new_optional_params) - assert new_tools[0]["function_declarations"][0]["parameters"]["properties"][ - "base_branch" - ] == { + assert new_tools[0]["function_declarations"][0]["parameters"]["properties"]["base_branch"] == { "anyOf": [{"type": "string", "nullable": True}] - }, f"Expected only anyOf field and its contents to be kept, but got {new_tools[0]['function_declarations'][0]['parameters']['properties']['base_branch']}" + }, ( + f"Expected only anyOf field and its contents to be kept, but got {new_tools[0]['function_declarations'][0]['parameters']['properties']['base_branch']}" + ) def test_vertex_ai_streaming_usage_calculation(): @@ -1328,7 +1307,7 @@ def test_vertex_ai_streaming_usage_calculation(): } # Test streaming chunk parsing - with patch.object(VertexGeminiConfig, "_calculate_usage") as mock_calculate_usage: + with patch.object(VertexGeminiConfig, "calculate_usage") as mock_calculate_usage: # Create a streaming chunk chunk = { "candidates": [{"content": {"parts": [{"text": "Hello"}]}}], @@ -1341,7 +1320,7 @@ def test_vertex_ai_streaming_usage_calculation(): ) iterator.chunk_parser(chunk) - # Verify _calculate_usage was called with correct parameters + # Verify calculate_usage was called with correct parameters mock_calculate_usage.assert_called_once_with(completion_response=chunk) # Test non-streaming response parsing @@ -1697,18 +1676,14 @@ def test_vertex_ai_usage_metadata_missing_token_count(): ], } usage_metadata = UsageMetadata(**usage_metadata) - result = v._calculate_usage(completion_response={"usageMetadata": usage_metadata}) + result = v.calculate_usage(completion_response={"usageMetadata": usage_metadata}) # Should not crash and should default missing tokenCount to 0 assert result.prompt_tokens == 57 assert result.completion_tokens == 74 assert result.total_tokens == 131 - assert ( - result.completion_tokens_details.text_tokens == 0 - ) # Default value for missing tokenCount - assert ( - result.completion_tokens_details.audio_tokens == 0 - ) # Default value for missing tokenCount + assert result.completion_tokens_details.text_tokens == 0 # Default value for missing tokenCount + assert result.completion_tokens_details.audio_tokens == 0 # Default value for missing tokenCount def test_vertex_ai_process_candidates_with_grounding_metadata(): @@ -1717,7 +1692,7 @@ def test_vertex_ai_process_candidates_with_grounding_metadata(): ) v = VertexGeminiConfig() - result = v._process_candidates( + result = v.process_candidates( _candidates=[ { "content": { @@ -1803,12 +1778,10 @@ def test_vertex_ai_process_candidates_with_grounding_metadata(): def test_set_stream_metadata_mirrors_non_streaming_safety_field_names(): - safety_ratings = [ - [{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}] - ] + safety_ratings = [[{"category": "HARM_CATEGORY_HATE_SPEECH", "probability": "NEGLIGIBLE"}]] model_response = ModelResponse() - VertexGeminiConfig._set_stream_metadata_on_response( + VertexGeminiConfig.set_stream_metadata_on_response( model_response=model_response, grounding_metadata=[], url_context_metadata=[], @@ -1914,7 +1887,7 @@ def test_vertex_ai_map_google_maps_tool_simple(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[{"googleMaps": {"enableWidget": "ENABLE_WIDGET"}}], optional_params=optional_params, ) @@ -1960,7 +1933,7 @@ def test_vertex_ai_map_google_maps_tool_with_location(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ { "googleMaps": { @@ -2405,39 +2378,34 @@ def test_is_gemini_3_or_newer(): ) # Gemini 3+ models - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-3-pro-preview") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-3-flash") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-3-pro") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-3.1-pro-preview") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-test-id-bla") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("test-id-bla") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-flash-latest") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-flash-lite-latest") == True - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-pro-latest") == True - assert ( - VertexGeminiConfig._is_gemini_3_or_newer("vertex_ai/gemini-3-pro-preview") - == True - ) - assert ( - VertexGeminiConfig._is_gemini_3_or_newer("gemini/gemini-3-pro-preview") == True - ) + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-3-pro-preview") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-3-flash") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-3-pro") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-3.1-pro-preview") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-test-id-bla") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("test-id-bla") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-flash-latest") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-flash-lite-latest") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-pro-latest") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("vertex_ai/gemini-3-pro-preview") == True + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini/gemini-3-pro-preview") == True # Gemini 2.5 and older models - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-2.5-pro") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-2.5-flash") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-2.0-flash") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-1.5-pro") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-pro") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini-flash") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-2.5-pro") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-2.5-flash") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-2.0-flash") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-1.5-pro") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-pro") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini-flash") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("4965075652664360960") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini/4965075652664360960") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini/ft-uuid") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemma-3-27b-it") == False - assert VertexGeminiConfig._is_gemini_3_or_newer("gemini/gemma-3-27b-it") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("4965075652664360960") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini/4965075652664360960") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini/ft-uuid") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemma-3-27b-it") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("gemini/gemma-3-27b-it") == False # Edge cases - assert VertexGeminiConfig._is_gemini_3_or_newer("") == False + assert VertexGeminiConfig.is_gemini_3_or_newer("") == False @pytest.mark.parametrize( @@ -2544,10 +2512,10 @@ def test_forward_gemini_function_call_id_is_gated_on_model_version_only(): VertexGeminiConfig, ) - assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3.5-flash") is True - assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-3-pro") is True - assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.5-flash") is False - assert VertexGeminiConfig._forward_gemini_function_call_id("gemini-2.0-flash") is False + assert VertexGeminiConfig.forward_gemini_function_call_id("gemini-3.5-flash") is True + assert VertexGeminiConfig.forward_gemini_function_call_id("gemini-3-pro") is True + assert VertexGeminiConfig.forward_gemini_function_call_id("gemini-2.5-flash") is False + assert VertexGeminiConfig.forward_gemini_function_call_id("gemini-2.0-flash") is False @pytest.mark.parametrize("custom_llm_provider", ["vertex_ai", "vertex_ai_beta", "gemini"]) @@ -2558,11 +2526,11 @@ def test_gemini_35_tool_calls_include_function_call_id(custom_llm_provider): side without the other would break strict tool-call matching. """ from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) tool_call_id = "call_50e7e0fe0989464a89f188eda443" - contents = _gemini_convert_messages_with_history( + contents = gemini_convert_messages_with_history( messages=_tool_call_messages(tool_call_id), model="gemini-3.5-flash", custom_llm_provider=custom_llm_provider, @@ -2575,10 +2543,10 @@ def test_gemini_35_tool_calls_include_function_call_id(custom_llm_provider): def test_gemini_25_tool_calls_omit_function_call_id(custom_llm_provider): """Regression: models older than Gemini 3 reject `id`, so the key must be absent entirely.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) - contents = _gemini_convert_messages_with_history( + contents = gemini_convert_messages_with_history( messages=_tool_call_messages("call_50e7e0fe0989464a89f188eda443"), model="gemini-2.5-flash", custom_llm_provider=custom_llm_provider, @@ -2599,15 +2567,15 @@ def test_vertex_ai_forwarded_function_call_id_strips_thought_signature_suffix(): Vertex now sees this code path for the first time, so the suffix has to be stripped here too. """ - from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, - ) from litellm.litellm_core_utils.prompt_templates.factory import ( THOUGHT_SIGNATURE_SEPARATOR, ) + from litellm.llms.vertex_ai.gemini.transformation import ( + gemini_convert_messages_with_history, + ) bare_id = "call_50e7e0fe0989464a89f188eda443" - contents = _gemini_convert_messages_with_history( + contents = gemini_convert_messages_with_history( messages=_tool_call_messages(f"{bare_id}{THOUGHT_SIGNATURE_SEPARATOR}sig123"), model="gemini-3.5-flash", custom_llm_provider="vertex_ai", @@ -2621,7 +2589,7 @@ def test_vertex_ai_forwarded_function_call_id_strips_thought_signature_suffix(): def test_tool_response_without_matching_tool_call_is_rejected(model): """An unpairable tool result must raise, not ship a functionResponse with no matching call.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -2644,7 +2612,7 @@ def test_tool_response_without_matching_tool_call_is_rejected(model): ] with pytest.raises(Exception, match="Missing corresponding tool call"): - _gemini_convert_messages_with_history( + gemini_convert_messages_with_history( messages=messages, model=model, custom_llm_provider="vertex_ai", @@ -2938,16 +2906,12 @@ def test_media_resolution_from_detail_parameter(): """Test that OpenAI's detail parameter is correctly mapped to media_resolution""" from litellm.llms.vertex_ai.gemini.transformation import ( _convert_detail_to_media_resolution_enum, - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) # Test detail -> media_resolution enum mapping - assert _convert_detail_to_media_resolution_enum("low") == { - "level": "MEDIA_RESOLUTION_LOW" - } - assert _convert_detail_to_media_resolution_enum("high") == { - "level": "MEDIA_RESOLUTION_HIGH" - } + assert _convert_detail_to_media_resolution_enum("low") == {"level": "MEDIA_RESOLUTION_LOW"} + assert _convert_detail_to_media_resolution_enum("high") == {"level": "MEDIA_RESOLUTION_HIGH"} assert _convert_detail_to_media_resolution_enum("auto") is None assert _convert_detail_to_media_resolution_enum(None) is None @@ -2966,9 +2930,7 @@ def test_media_resolution_from_detail_parameter(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") # Verify media_resolution is set at the Part level (not inside inline_data) assert len(contents) == 1 @@ -2989,7 +2951,7 @@ def test_media_resolution_from_detail_parameter(): def test_media_resolution_low_detail(): """Test that detail='low' maps to media_resolution enum with MEDIA_RESOLUTION_LOW""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) # Using a minimal valid base64-encoded 1x1 PNG @@ -3006,9 +2968,7 @@ def test_media_resolution_low_detail(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") # Find the part with inline_data image_part = None @@ -3026,7 +2986,7 @@ def test_media_resolution_low_detail(): def test_media_resolution_auto_detail(): """Test that detail='auto' or None doesn't set media_resolution""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) # Using a minimal valid base64-encoded 1x1 PNG @@ -3045,7 +3005,7 @@ def test_media_resolution_auto_detail(): } ] - contents = _gemini_convert_messages_with_history(messages=messages_auto) + contents = gemini_convert_messages_with_history(messages=messages_auto) # Find the part with inline_data image_part = None for part in contents[0]["parts"]: @@ -3065,7 +3025,7 @@ def test_media_resolution_auto_detail(): } ] - contents = _gemini_convert_messages_with_history(messages=messages_none) + contents = gemini_convert_messages_with_history(messages=messages_none) # Find the part with inline_data image_part = None for part in contents[0]["parts"]: @@ -3081,7 +3041,7 @@ def test_media_resolution_auto_detail(): def test_media_resolution_per_part(): """Test that different images can have different media_resolution values""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) # Using minimal valid base64-encoded 1x1 PNGs @@ -3105,9 +3065,7 @@ def test_media_resolution_per_part(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") # Should have one content with multiple parts assert len(contents) == 1 @@ -3131,7 +3089,7 @@ def test_media_resolution_per_part(): def test_media_resolution_only_for_gemini_3_models(): """Ensure media_resolution is not added for non-Gemini 3 models.""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) base64_image = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==" @@ -3150,9 +3108,7 @@ def test_media_resolution_only_for_gemini_3_models(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-pro" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-2.5-pro") image_part = None for part in contents[0]["parts"]: if "inline_data" in part: @@ -3482,7 +3438,7 @@ def test_vertex_ai_multiple_tool_types_separate_objects(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ {"enterpriseWebSearch": {}}, {"url_context": {}}, @@ -3538,7 +3494,7 @@ def test_vertex_ai_function_declarations_with_other_tools_separate(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ { "type": "function", @@ -3585,9 +3541,7 @@ def test_vertex_ai_single_tool_type_still_works(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( - value=[{"code_execution": {}}], optional_params=optional_params - ) + tools = v.map_function(value=[{"code_execution": {}}], optional_params=optional_params) assert len(tools) == 1 assert "code_execution" in tools[0] @@ -3608,7 +3562,7 @@ def test_vertex_ai_mixed_search_and_function_tools_drops_search(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ {"enterpriseWebSearch": {}}, {"urlContext": {}}, @@ -3640,7 +3594,7 @@ def test_vertex_ai_mixed_google_search_and_function_tools_drops_search(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ {"googleSearch": {}}, { @@ -3663,7 +3617,7 @@ def test_vertex_ai_search_tools_only_no_drop(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ {"enterpriseWebSearch": {}}, {"urlContext": {}}, @@ -3685,7 +3639,7 @@ def test_vertex_ai_function_tools_with_code_execution_preserved(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ {"code_execution": {}}, { @@ -3712,7 +3666,7 @@ def test_vertex_ai_gemini3_tool_combination_no_drop(): v = VertexGeminiConfig() optional_params = {"include_server_side_tool_invocations": True} - tools = v._map_function( + tools = v.map_function( value=[ {"enterpriseWebSearch": {}}, {"urlContext": {}}, @@ -3862,17 +3816,11 @@ def test_vertex_ai_openai_web_search_tool_transformation(): optional_params = {} # Test web_search transformation - tools = v._map_function( - value=[{"type": "web_search"}], optional_params=optional_params - ) + tools = v.map_function(value=[{"type": "web_search"}], optional_params=optional_params) assert len(tools) == 1, f"Expected 1 Tool object, got {len(tools)}" - assert ( - "googleSearch" in tools[0] - ), f"Expected googleSearch in tool, got {tools[0].keys()}" - assert ( - tools[0]["googleSearch"] == {} - ), f"Expected empty googleSearch config, got {tools[0]['googleSearch']}" + assert "googleSearch" in tools[0], f"Expected googleSearch in tool, got {tools[0].keys()}" + assert tools[0]["googleSearch"] == {}, f"Expected empty googleSearch config, got {tools[0]['googleSearch']}" def test_vertex_ai_openai_web_search_preview_tool_transformation(): @@ -3889,17 +3837,11 @@ def test_vertex_ai_openai_web_search_preview_tool_transformation(): optional_params = {} # Test web_search_preview transformation - tools = v._map_function( - value=[{"type": "web_search_preview"}], optional_params=optional_params - ) + tools = v.map_function(value=[{"type": "web_search_preview"}], optional_params=optional_params) assert len(tools) == 1, f"Expected 1 Tool object, got {len(tools)}" - assert ( - "googleSearch" in tools[0] - ), f"Expected googleSearch in tool, got {tools[0].keys()}" - assert ( - tools[0]["googleSearch"] == {} - ), f"Expected empty googleSearch config, got {tools[0]['googleSearch']}" + assert "googleSearch" in tools[0], f"Expected googleSearch in tool, got {tools[0].keys()}" + assert tools[0]["googleSearch"] == {}, f"Expected empty googleSearch config, got {tools[0]['googleSearch']}" def test_vertex_ai_openai_web_search_with_function_tools(): @@ -3922,7 +3864,7 @@ def test_vertex_ai_openai_web_search_with_function_tools(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ {"type": "web_search"}, { @@ -3968,7 +3910,7 @@ def test_vertex_ai_multiple_function_declarations_grouped(): v = VertexGeminiConfig() optional_params = {} - tools = v._map_function( + tools = v.map_function( value=[ { "type": "function", @@ -4009,7 +3951,7 @@ def test_gemini_3_flash_preview_token_usage_fallback(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) assert result.completion_tokens == 509 assert result.prompt_tokens == 2145 @@ -4033,7 +3975,7 @@ def test_gemini_no_reasoning_fallback(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) assert result.completion_tokens == 264 assert result.completion_tokens_details is not None @@ -4059,7 +4001,7 @@ def test_gemini_token_usage_standard_response(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) assert result.completion_tokens == 50 assert result.completion_tokens_details.text_tokens == 40 @@ -4103,7 +4045,7 @@ def test_gemini_image_gen_usage_metadata_prompt_vs_completion_separation(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) # Verify basic token counts assert result.prompt_tokens == 101 @@ -4111,32 +4053,26 @@ def test_gemini_image_gen_usage_metadata_prompt_vs_completion_separation(): assert result.total_tokens == 1391 # CRITICAL: Prompt tokens details should show NO image tokens (text-only input) - assert ( - result.prompt_tokens_details.text_tokens == 101 - ), "Prompt text tokens should be 101" - assert ( - result.prompt_tokens_details.image_tokens is None - ), "Prompt image tokens should be None (text-only input, no images in prompt)" - assert ( - result.prompt_tokens_details.audio_tokens is None - ), "Prompt audio tokens should be None" + assert result.prompt_tokens_details.text_tokens == 101, "Prompt text tokens should be 101" + assert result.prompt_tokens_details.image_tokens is None, ( + "Prompt image tokens should be None (text-only input, no images in prompt)" + ) + assert result.prompt_tokens_details.audio_tokens is None, "Prompt audio tokens should be None" # Completion tokens details should show the generated image tokens - assert ( - result.completion_tokens_details.image_tokens == 1290 - ), "Completion image tokens should be 1290 (generated image)" + assert result.completion_tokens_details.image_tokens == 1290, ( + "Completion image tokens should be 1290 (generated image)" + ) # Verify text_tokens is auto-calculated for completion # candidatesTokenCount (1290) - image_tokens (1290) = 0 - assert ( - result.completion_tokens_details.text_tokens == 0 - ), "Completion text tokens should be 0 (image-only response)" + assert result.completion_tokens_details.text_tokens == 0, "Completion text tokens should be 0 (image-only response)" def test_file_object_detail_parameter(): """Test that detail parameter works for type: file objects (Issue #19026)""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -4156,9 +4092,7 @@ def test_file_object_detail_parameter(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") # Verify media_resolution is set for file objects assert len(contents) == 1 @@ -4172,16 +4106,14 @@ def test_file_object_detail_parameter(): break assert file_part is not None, "File part should exist" - assert ( - "media_resolution" in file_part - ), "media_resolution should be set for file objects" + assert "media_resolution" in file_part, "media_resolution should be set for file objects" assert file_part["media_resolution"] == {"level": "MEDIA_RESOLUTION_LOW"} def test_video_metadata_fps(): """Test fps parameter in video_metadata (Issue #19026)""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -4201,9 +4133,7 @@ def test_video_metadata_fps(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") # Find the file part file_part = None @@ -4220,7 +4150,7 @@ def test_video_metadata_fps(): def test_video_metadata_complete(): """Test all video_metadata fields: fps, start_offset, end_offset (Issue #19026)""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -4244,9 +4174,7 @@ def test_video_metadata_complete(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") # Find the file part file_part = None @@ -4268,7 +4196,7 @@ def test_video_metadata_complete(): def test_detail_and_video_metadata_combined(): """Test using both detail and video_metadata together (Issue #19026)""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -4289,9 +4217,7 @@ def test_detail_and_video_metadata_combined(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") # Find the file part file_part = None @@ -4311,22 +4237,14 @@ def test_new_detail_levels(): """Test new detail levels: medium and ultra_high (Issue #19026)""" from litellm.llms.vertex_ai.gemini.transformation import ( _convert_detail_to_media_resolution_enum, - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) # Test mapping function - assert _convert_detail_to_media_resolution_enum("low") == { - "level": "MEDIA_RESOLUTION_LOW" - } - assert _convert_detail_to_media_resolution_enum("medium") == { - "level": "MEDIA_RESOLUTION_MEDIUM" - } - assert _convert_detail_to_media_resolution_enum("high") == { - "level": "MEDIA_RESOLUTION_HIGH" - } - assert _convert_detail_to_media_resolution_enum("ultra_high") == { - "level": "MEDIA_RESOLUTION_ULTRA_HIGH" - } + assert _convert_detail_to_media_resolution_enum("low") == {"level": "MEDIA_RESOLUTION_LOW"} + assert _convert_detail_to_media_resolution_enum("medium") == {"level": "MEDIA_RESOLUTION_MEDIUM"} + assert _convert_detail_to_media_resolution_enum("high") == {"level": "MEDIA_RESOLUTION_HIGH"} + assert _convert_detail_to_media_resolution_enum("ultra_high") == {"level": "MEDIA_RESOLUTION_ULTRA_HIGH"} # Test with actual message transformation messages = [ @@ -4345,9 +4263,7 @@ def test_new_detail_levels(): } ] - contents = _gemini_convert_messages_with_history( - messages=messages, model="gemini-3-pro-preview" - ) + contents = gemini_convert_messages_with_history(messages=messages, model="gemini-3-pro-preview") file_part = None for part in contents[0]["parts"]: @@ -4362,7 +4278,7 @@ def test_new_detail_levels(): def test_video_metadata_supported_for_all_gemini_models(): """Test that video_metadata is applied for all Gemini models (Issue #25474)""" from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) messages = [ @@ -4388,7 +4304,7 @@ def test_video_metadata_supported_for_all_gemini_models(): "gemini-2.5-pro", "gemini-3-pro-preview", ]: - contents = _gemini_convert_messages_with_history(messages=messages, model=model) + contents = gemini_convert_messages_with_history(messages=messages, model=model) file_part = None for part in contents[0]["parts"]: @@ -4397,25 +4313,19 @@ def test_video_metadata_supported_for_all_gemini_models(): break assert file_part is not None, f"{model}: file part should exist" - assert ( - "video_metadata" in file_part - ), f"{model}: video_metadata should be present" + assert "video_metadata" in file_part, f"{model}: video_metadata should be present" assert file_part["video_metadata"]["fps"] == 5, f"{model}: fps should be 5" # Per-part media_resolution is Gemini 3+ only; 2.x uses generation_config global for model in ["gemini-3-pro-preview"]: - contents = _gemini_convert_messages_with_history(messages=messages, model=model) + contents = gemini_convert_messages_with_history(messages=messages, model=model) file_part = next(p for p in contents[0]["parts"] if "file_data" in p) - assert ( - "media_resolution" in file_part - ), f"{model}: media_resolution should be present" + assert "media_resolution" in file_part, f"{model}: media_resolution should be present" for model in ["gemini-1.5-pro", "gemini-2.5-flash", "gemini-2.5-pro"]: - contents = _gemini_convert_messages_with_history(messages=messages, model=model) + contents = gemini_convert_messages_with_history(messages=messages, model=model) file_part = next(p for p in contents[0]["parts"] if "file_data" in p) - assert ( - "media_resolution" not in file_part - ), f"{model}: per-part media_resolution should not be set" + assert "media_resolution" not in file_part, f"{model}: per-part media_resolution should not be set" def test_chunk_parser_handles_prompt_feedback_block(): @@ -4741,15 +4651,11 @@ def test_vertex_ai_web_search_options_parameter(): # When web_search_options is present, it should be mapped to a tool web_search_options = {} - _tools = v._map_web_search_options(web_search_options) + _tools = v.map_web_search_options(web_search_options) # Verify the tool is a googleSearch tool - assert ( - "googleSearch" in _tools - ), f"Expected googleSearch in tool, got {_tools.keys()}" - assert ( - _tools["googleSearch"] == {} - ), f"Expected empty googleSearch config, got {_tools['googleSearch']}" + assert "googleSearch" in _tools, f"Expected googleSearch in tool, got {_tools.keys()}" + assert _tools["googleSearch"] == {}, f"Expected empty googleSearch config, got {_tools['googleSearch']}" def test_vertex_ai_web_search_options_in_map_openai_params(): @@ -4777,12 +4683,10 @@ def test_vertex_ai_web_search_options_in_map_openai_params(): # Call the transformation that happens in map_openai_params # Lines 1075-1079 in vertex_and_google_ai_studio_gemini.py (after fix) web_search_value = optional_params.get("web_search_options") - if isinstance( - web_search_value, dict - ): # Fixed: removed 'value and' check to support empty dicts - _tools = v._map_web_search_options(web_search_value) + if isinstance(web_search_value, dict): # Fixed: removed 'value and' check to support empty dicts + _tools = v.map_web_search_options(web_search_value) # Simulate _add_tools_to_optional_params - optional_params = v._add_tools_to_optional_params(optional_params, [_tools]) + optional_params = v.add_tools_to_optional_params(optional_params, [_tools]) # Remove web_search_options as it's been transformed optional_params.pop("web_search_options", None) @@ -4875,7 +4779,7 @@ def test_vertex_ai_usage_metadata_with_video_tokens_in_prompt(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) # Verify basic token counts assert result.prompt_tokens == 10449 @@ -4927,21 +4831,15 @@ def test_vertex_ai_usage_metadata_with_video_tokens_in_candidates(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) assert result.completion_tokens == 10330 assert result.completion_tokens_details is not None - assert ( - result.completion_tokens_details.video_tokens == 10240 - ), "Completion video tokens should be 10240" - assert ( - result.completion_tokens_details.text_tokens == 90 - ), "Completion text tokens should be 90" + assert result.completion_tokens_details.video_tokens == 10240, "Completion video tokens should be 10240" + assert result.completion_tokens_details.text_tokens == 90, "Completion text tokens should be 90" # Verify prompt side has no video tokens - assert ( - result.prompt_tokens_details.video_tokens is None - ), "Prompt video tokens should be None (text-only input)" + assert result.prompt_tokens_details.video_tokens is None, "Prompt video tokens should be None (text-only input)" def test_vertex_ai_usage_metadata_video_tokens_auto_calculated_text(): @@ -4963,13 +4861,13 @@ def test_vertex_ai_usage_metadata_video_tokens_auto_calculated_text(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) assert result.completion_tokens_details.video_tokens == 10240 # text = 10330 - 10240 = 90 - assert ( - result.completion_tokens_details.text_tokens == 90 - ), "text_tokens should be auto-calculated as candidatesTokenCount - video_tokens" + assert result.completion_tokens_details.text_tokens == 90, ( + "text_tokens should be auto-calculated as candidatesTokenCount - video_tokens" + ) def test_vertex_ai_usage_metadata_video_tokens_with_caching(): @@ -4997,12 +4895,12 @@ def test_vertex_ai_usage_metadata_video_tokens_with_caching(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) # video tokens should be reduced by cached amount: 10240 - 5120 = 5120 - assert ( - result.prompt_tokens_details.video_tokens == 5120 - ), "Prompt video tokens should be 10240 - 5120 (cached) = 5120" + assert result.prompt_tokens_details.video_tokens == 5120, ( + "Prompt video tokens should be 10240 - 5120 (cached) = 5120" + ) assert result.prompt_tokens_details.text_tokens == 9 assert result.prompt_tokens_details.audio_tokens == 200 @@ -5039,7 +4937,7 @@ def test_vertex_ai_usage_metadata_with_document_tokens_in_prompt(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) # Verify basic token counts assert result.prompt_tokens == 782 @@ -5081,13 +4979,11 @@ def test_vertex_ai_usage_metadata_with_document_tokens_cached(): } completion_response = {"usageMetadata": usage_metadata_dict} - result = v._calculate_usage(completion_response=completion_response) + result = v.calculate_usage(completion_response=completion_response) # DOCUMENT cached tokens map to cached_text_tokens, so: # text_tokens = (8 TEXT + 774 DOCUMENT) - 400 cached = 382 - assert ( - result.prompt_tokens_details.text_tokens == 382 - ), "text_tokens should be (8 + 774) - 400 cached = 382" + assert result.prompt_tokens_details.text_tokens == 382, "text_tokens should be (8 + 774) - 400 cached = 382" assert result.prompt_tokens_details.cached_tokens == 400 @@ -5801,7 +5697,7 @@ def test_process_candidates_merges_thought_signatures_and_server_side_tools(): ] model_response = ModelResponse() - VertexGeminiConfig._process_candidates( + VertexGeminiConfig.process_candidates( _candidates=candidates, model_response=model_response, standard_optional_params={}, @@ -5954,16 +5850,16 @@ def test_calculate_web_search_requests_counts_unique_queries(): duplicates_in_one_item: Final = [ {"webSearchQueries": ["euro 2024 winner", "euro 2024 winner", "spain england final", ""]} ] - assert VertexGeminiConfig._calculate_web_search_requests(duplicates_in_one_item) == 2 + assert VertexGeminiConfig.calculate_web_search_requests(duplicates_in_one_item) == 2 duplicates_across_items: Final = [ {"webSearchQueries": ["euro 2024 winner"]}, {"webSearchQueries": ["euro 2024 winner", "spain england final"]}, ] - assert VertexGeminiConfig._calculate_web_search_requests(duplicates_across_items) == 2 + assert VertexGeminiConfig.calculate_web_search_requests(duplicates_across_items) == 2 - assert VertexGeminiConfig._calculate_web_search_requests([]) is None - assert VertexGeminiConfig._calculate_web_search_requests([{"webSearchQueries": ["", ""]}]) is None + assert VertexGeminiConfig.calculate_web_search_requests([]) is None + assert VertexGeminiConfig.calculate_web_search_requests([{"webSearchQueries": ["", ""]}]) is None @pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai"]) @@ -6056,7 +5952,7 @@ def test_generate_content_transform_uses_reported_model_version(): import httpx body = {**_generate_content_body(), "modelVersion": "gemini-x-served"} - response: Final = VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response( + response: Final = VertexGeminiConfig().transform_google_generate_content_to_openai_model_response( completion_response=body, model_response=ModelResponse(), model="gemini-x", @@ -6070,7 +5966,7 @@ def test_generate_content_transform_uses_reported_model_version(): def test_generate_content_transform_falls_back_to_requested_model(): import httpx - response: Final = VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response( + response: Final = VertexGeminiConfig().transform_google_generate_content_to_openai_model_response( completion_response=_generate_content_body(), model_response=ModelResponse(), model="gemini-x", @@ -6127,7 +6023,7 @@ def test_generate_content_transform_strips_version_suffix_from_model_version(): import httpx body: Final = {**_generate_content_body(), "modelVersion": "gemini-3.8-flash-001@default"} - response: Final = VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response( + response: Final = VertexGeminiConfig().transform_google_generate_content_to_openai_model_response( completion_response=body, model_response=ModelResponse(), model="gemini-3.8-flash", @@ -6176,7 +6072,7 @@ def test_gemini_candidate_with_finish_reason_no_content_chat_completion(): raw_response = MagicMock() raw_response.headers = {} - resp = config._transform_google_generate_content_to_openai_model_response( + resp = config.transform_google_generate_content_to_openai_model_response( completion_response=completion_response, model_response=model_response, model="gemini-2.5-flash-image", @@ -6208,7 +6104,7 @@ def test_gemini_candidate_with_finish_reason_no_content_anthropic_messages(): "totalTokenCount": 19, }, } - resp = config._transform_google_generate_content_to_openai_model_response( + resp = config.transform_google_generate_content_to_openai_model_response( completion_response=completion_response, model_response=ModelResponse(), model="gemini-2.5-flash-image", @@ -6244,7 +6140,7 @@ def test_gemini_candidate_with_finish_reason_no_content_responses_api(): "totalTokenCount": 19, }, } - resp = config._transform_google_generate_content_to_openai_model_response( + resp = config.transform_google_generate_content_to_openai_model_response( completion_response=completion_response, model_response=ModelResponse(), model="gemini-2.5-flash-image", @@ -6275,7 +6171,7 @@ def test_gemini_candidate_other_finish_reasons_no_content(): "candidates": [{"finishReason": "MAX_TOKENS", "index": 0}], "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 50, "totalTokenCount": 60}, } - resp_length = config._transform_google_generate_content_to_openai_model_response( + resp_length = config.transform_google_generate_content_to_openai_model_response( completion_response=max_tokens_response, model_response=ModelResponse(), model="gemini-2.5-flash", @@ -6344,7 +6240,7 @@ def test_gemini_multi_candidate_messages_do_not_share_state(): "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 20, "totalTokenCount": 30}, } - resp: Final = config._transform_google_generate_content_to_openai_model_response( + resp: Final = config.transform_google_generate_content_to_openai_model_response( completion_response=completion_response, model_response=ModelResponse(), model="gemini-2.5-flash", @@ -6553,7 +6449,7 @@ def test_round_trip_thought_signature_in_conversation(): {"role": "user", "content": "How are you?"}, ] - gemini_contents = _gemini_convert_messages_with_history(messages) + gemini_contents = gemini_convert_messages_with_history(messages) # Find the assistant (model) message model_message = None @@ -6583,7 +6479,7 @@ def test_round_trip_without_thought_signature_still_works(): {"role": "user", "content": "How are you?"}, ] - gemini_contents = _gemini_convert_messages_with_history(messages) + gemini_contents = gemini_convert_messages_with_history(messages) # Find the assistant (model) message model_message = None diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py index b6e7f20f159..23b14f84016 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_gemini_unbound_local_error.py @@ -20,7 +20,7 @@ def test_process_candidates_unbound_local_error_fix(): # Execution try: - VertexGeminiConfig._process_candidates( + VertexGeminiConfig.process_candidates( _candidates=candidates, model_response=model_response, standard_optional_params={}, diff --git a/tests/unit/llms/vertex_ai/test_bge_response_transformation.py b/tests/unit/llms/vertex_ai/test_bge_response_transformation.py index 9e960570036..8c6734829bb 100644 --- a/tests/unit/llms/vertex_ai/test_bge_response_transformation.py +++ b/tests/unit/llms/vertex_ai/test_bge_response_transformation.py @@ -18,7 +18,7 @@ def test_is_bge_model_detection(): Test BGE model detection for post-provider-split patterns. After main.py splits the provider, model strings are passed without the provider prefix. - Model name transformation (bge/ -> numeric ID) is handled in common_utils._get_vertex_url(). + Model name transformation (bge/ -> numeric ID) is handled in common_utils.get_vertex_url(). """ # Should detect BGE models (after provider split) assert VertexBGEConfig.is_bge_model("bge-small-en-v1.5") is True diff --git a/tests/unit/llms/vertex_ai/test_common_utils.py b/tests/unit/llms/vertex_ai/test_common_utils.py index 19450ffe140..3c8a4a4a34d 100644 --- a/tests/unit/llms/vertex_ai/test_common_utils.py +++ b/tests/unit/llms/vertex_ai/test_common_utils.py @@ -14,14 +14,14 @@ from typing_extensions import ReadOnly, TypedDict import litellm from litellm import Router from litellm.constants import INITIAL_RETRY_DELAY, MAX_RETRY_DELAY -from litellm.llms.vertex_ai.common_utils import _get_gemini_url, get_vertex_base_url +from litellm.llms.vertex_ai.common_utils import get_gemini_url, get_vertex_base_url from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from tests._vcr_conftest_common import install_live_call_probe, record_vcr_outcome CONFIG_PATH: Final = Path(__file__).parent / "google_genai_proxy_test_config.yaml" GEMINI_DEPLOYMENT: Final = "gemini-3.5-flash-lite" VERTEX_DEPLOYMENT: Final = "vertex-gemini-3.5-flash-lite" -GEMINI_GENERATE_CONTENT_URL: Final = _get_gemini_url(mode="chat", model=GEMINI_DEPLOYMENT, stream=False)[0] +GEMINI_GENERATE_CONTENT_URL: Final = get_gemini_url(mode="chat", model=GEMINI_DEPLOYMENT, stream=False)[0] VERTEX_GLOBAL_BASE_URL: Final = "https://aiplatform.googleapis.com" RESOURCE_EXHAUSTED: Final = { "error": {"code": 429, "message": "Resource exhausted. Please try again later.", "status": "RESOURCE_EXHAUSTED"} diff --git a/tests/unit/llms/vertex_ai/test_vertex.py b/tests/unit/llms/vertex_ai/test_vertex.py index fe986371b73..94dc7dc861d 100644 --- a/tests/unit/llms/vertex_ai/test_vertex.py +++ b/tests/unit/llms/vertex_ai/test_vertex.py @@ -104,7 +104,7 @@ def test_completion_pydantic_obj_2(): def test_build_vertex_schema(): import json - from litellm.llms.vertex_ai.common_utils import _build_vertex_schema + from litellm.llms.vertex_ai.common_utils import build_vertex_schema schema = { "type": "object", @@ -122,7 +122,7 @@ def test_build_vertex_schema(): "required": ["recipes"], } - new_schema = _build_vertex_schema(schema) + new_schema = build_vertex_schema(schema) print(f"new_schema: {new_schema}") assert new_schema["type"] == schema["type"] assert new_schema["properties"] == schema["properties"] @@ -1223,12 +1223,10 @@ def test_process_gemini_media(): } ] from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) - converted = _gemini_convert_messages_with_history( - messages=image_message, model="gemini-2.5-flash" - ) + converted = gemini_convert_messages_with_history(messages=image_message, model="gemini-2.5-flash") assert converted[0]["parts"][0]["file_data"] == FileDataType( mime_type="image/png", file_uri="gs://bucket/image-without-extension" ) @@ -1329,9 +1327,9 @@ def test_vertex_embedding_url(model, expected_url): When a fine-tuned embedding model is used, the URL is different from the standard one. """ - from litellm.llms.vertex_ai.common_utils import _get_vertex_url + from litellm.llms.vertex_ai.common_utils import get_vertex_url - url, endpoint = _get_vertex_url( + url, endpoint = get_vertex_url( mode="embedding", model=model, stream=False, @@ -1506,7 +1504,7 @@ def test_vertex_parallel_tool_calls_false_single_tool(): assert "tools" in optional_params -from litellm.llms.vertex_ai.gemini.transformation import _transform_request_body +from litellm.llms.vertex_ai.gemini.transformation import transform_request_body def test_system_prompt_only_adds_blank_user_message(): @@ -1516,7 +1514,7 @@ def test_system_prompt_only_adds_blank_user_message(): Relevant Issue - https://github.com/BerriAI/litellm/issues/13769 """ SYSTEM_INSTRUCTION = "System instructions for the model" - data = _transform_request_body( + data = transform_request_body( messages=[{"role": "system", "content": SYSTEM_INSTRUCTION}], model="gemini-2.5-flash", optional_params={}, diff --git a/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py b/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py index 3bc66d017a0..8da95f839b9 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py +++ b/tests/unit/llms/vertex_ai/test_vertex_ai_batch_transformation.py @@ -56,14 +56,10 @@ def test_vertex_ai_cancel_batch(): "state": "JOB_STATE_CANCELLING", "createTime": "2024-03-17T10:00:00.000000Z", "inputConfig": {"gcsSource": {"uris": ["gs://test-bucket/input.jsonl"]}}, - "outputConfig": { - "gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"} - }, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"}}, } - with patch( - "litellm.llms.vertex_ai.batches.handler._get_httpx_client" - ) as mock_client: + with patch("litellm.llms.vertex_ai.batches.handler.get_httpx_client") as mock_client: mock_client.return_value.post.return_value = mock_response mock_client.return_value.get.return_value = mock_response @@ -101,14 +97,10 @@ def test_vertex_ai_cancel_batch_encodes_batch_id(): "state": "JOB_STATE_CANCELLING", "createTime": "2024-03-17T10:00:00.000000Z", "inputConfig": {"gcsSource": {"uris": ["gs://test-bucket/input.jsonl"]}}, - "outputConfig": { - "gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"} - }, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"}}, } - with patch( - "litellm.llms.vertex_ai.batches.handler._get_httpx_client" - ) as mock_client: + with patch("litellm.llms.vertex_ai.batches.handler.get_httpx_client") as mock_client: mock_client.return_value.post.return_value = mock_response mock_client.return_value.get.return_value = mock_response @@ -154,14 +146,10 @@ def test_vertex_ai_cancel_batch_custom_proxy_retrieve_url(): "state": "JOB_STATE_CANCELLING", "createTime": "2024-03-17T10:00:00.000000Z", "inputConfig": {"gcsSource": {"uris": ["gs://test-bucket/input.jsonl"]}}, - "outputConfig": { - "gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"} - }, + "outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"}}, } - with patch( - "litellm.llms.vertex_ai.batches.handler._get_httpx_client" - ) as mock_client: + with patch("litellm.llms.vertex_ai.batches.handler.get_httpx_client") as mock_client: mock_client.return_value.post.return_value = mock_response mock_client.return_value.get.return_value = mock_response diff --git a/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py index 86eb26a4c15..8407007e1d4 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -3,13 +3,11 @@ from unittest.mock import patch import pytest from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH - - from litellm.llms.vertex_ai.common_utils import ( - _get_vertex_url, convert_anyof_null_to_nullable, get_vertex_location_from_url, get_vertex_project_id_from_url, + get_vertex_url, pop_vertex_request_labels, set_schema_property_ordering, supports_response_json_schema, @@ -213,7 +211,7 @@ def test_set_schema_property_ordering_skips_non_dict_property_values(): def test_build_vertex_schema(): """Test build_vertex_schema with a sample schema""" - from litellm.llms.vertex_ai.common_utils import _build_vertex_schema + from litellm.llms.vertex_ai.common_utils import build_vertex_schema parameters = { "properties": { @@ -295,7 +293,7 @@ def test_build_vertex_schema(): "type": "object", } - assert _build_vertex_schema(parameters) == expected_output + assert build_vertex_schema(parameters) == expected_output def test_process_items_with_excessive_nesting(): @@ -363,7 +361,7 @@ def test_build_vertex_schema_array_branch_missing_items_in_anyof(): end up with synthesized `items: {"type": "object"}` after the schema transform — Vertex returns INVALID_ARGUMENT otherwise. """ - from litellm.llms.vertex_ai.common_utils import _build_vertex_schema + from litellm.llms.vertex_ai.common_utils import build_vertex_schema parameters = { "properties": { @@ -378,14 +376,12 @@ def test_build_vertex_schema_array_branch_missing_items_in_anyof(): "type": "object", } - result = _build_vertex_schema(parameters) + result = build_vertex_schema(parameters) callbacks_anyof = result["properties"]["callbacks"]["anyOf"] array_branches = [b for b in callbacks_anyof if b.get("type") == "array"] assert array_branches, "expected an array branch to remain after transform" for branch in array_branches: - assert branch.get("items") == { - "type": "object" - }, f"array branch must have items synthesized; got {branch}" + assert branch.get("items") == {"type": "object"}, f"array branch must have items synthesized; got {branch}" def test_vertex_ai_complex_response_schema(): @@ -622,7 +618,7 @@ def test_vertex_ai_complex_response_schema(): ) def test_get_vertex_url_global_region(stream, expected_endpoint_suffix): """ - Test _get_vertex_url when vertex_location is 'global' for chat mode. + Test get_vertex_url when vertex_location is 'global' for chat mode. """ mode = "chat" model = "gemini-1.5-pro-preview-0409" @@ -636,7 +632,7 @@ def test_get_vertex_url_global_region(stream, expected_endpoint_suffix): "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, endpoint = _get_vertex_url( + url, endpoint = get_vertex_url( mode=mode, model=model, stream=stream, @@ -1149,7 +1145,7 @@ def test_get_token_url(): vertex_ai_location = "us-central1" vertex_credentials = "" - _, url = vertex_llm._get_token_and_url( + _, url = vertex_llm.get_token_and_url( auth_header=None, vertex_project=vertex_ai_project, vertex_location=vertex_ai_location, @@ -1164,7 +1160,7 @@ def test_get_token_url(): print("url=", url) - _, url = vertex_llm._get_token_and_url( + _, url = vertex_llm.get_token_and_url( auth_header=None, vertex_project=vertex_ai_project, vertex_location=vertex_ai_location, @@ -1683,7 +1679,7 @@ def test_vertex_ai_google_gemini_not_detected_as_gemma_maas(): def test_build_vertex_schema_empty_properties(): """ - Test _build_vertex_schema handles empty properties objects correctly. + Test build_vertex_schema handles empty properties objects correctly. This test verifies the fix for the issue where Gemini rejects schemas with empty properties objects like {"properties": {}, "type": "object"}. @@ -1694,7 +1690,7 @@ def test_build_vertex_schema_empty_properties(): The fix removes empty properties objects and their associated type/required fields. """ - from litellm.llms.vertex_ai.common_utils import _build_vertex_schema + from litellm.llms.vertex_ai.common_utils import build_vertex_schema # Input: Schema with empty properties (the problematic case from real request) input_schema = { @@ -1727,40 +1723,28 @@ def test_build_vertex_schema_empty_properties(): } # Apply the transformation - result = _build_vertex_schema(input_schema) + result = build_vertex_schema(input_schema) # Verify the transformation removed empty properties # Navigate to the go_back schema - go_back_schema = result["properties"]["action"]["items"]["anyOf"][0]["properties"][ - "go_back" - ] + go_back_schema = result["properties"]["action"]["items"]["anyOf"][0]["properties"]["go_back"] # Verify empty properties was removed assert "properties" not in go_back_schema, "Empty properties should be removed" # Verify type is kept as object (Gemini requires type: object even without properties) - assert ( - go_back_schema.get("type") == "object" - ), "Type should be kept as object when properties is empty" + assert go_back_schema.get("type") == "object", "Type should be kept as object when properties is empty" # Verify required was also removed - assert ( - "required" not in go_back_schema - ), "Required should be removed when properties is empty" + assert "required" not in go_back_schema, "Required should be removed when properties is empty" # Verify description is preserved - assert ( - go_back_schema.get("description") == "Go back" - ), "Description should be preserved" + assert go_back_schema.get("description") == "Go back", "Description should be preserved" # Verify parent schema still has proper structure parent_schema = result["properties"]["action"]["items"]["anyOf"][0] - assert ( - parent_schema["type"] == "object" - ), "Parent schema should still have object type" - assert ( - "go_back" in parent_schema["properties"] - ), "go_back should still be in parent properties" + assert parent_schema["type"] == "object", "Parent schema should still have object type" + assert "go_back" in parent_schema["properties"], "go_back should still be in parent properties" def test_add_object_type_schema_with_no_properties_and_no_type(): diff --git a/tests/unit/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py b/tests/unit/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py index 55493d47f3d..3906528ae2e 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py +++ b/tests/unit/llms/vertex_ai/test_vertex_gemini_gcs_uri_mime.py @@ -81,7 +81,7 @@ def test_process_gemini_media_rejects_gcs_metadata_mime_not_supported_by_gemini( def test_file_block_uses_mime_type_alias_for_extensionless_gcs(): from litellm.llms.vertex_ai.gemini.transformation import ( - _gemini_convert_messages_with_history, + gemini_convert_messages_with_history, ) from litellm.types.llms.vertex_ai import FileDataType @@ -99,9 +99,7 @@ def test_file_block_uses_mime_type_alias_for_extensionless_gcs(): ], } ] - converted = _gemini_convert_messages_with_history( - messages=messages, model="gemini-2.5-flash" - ) + converted = gemini_convert_messages_with_history(messages=messages, model="gemini-2.5-flash") assert converted[0]["parts"][0]["file_data"] == FileDataType( mime_type="application/pdf", file_uri="gs://bucket/no-extension-object" ) diff --git a/tests/unit/llms/vertex_ai/test_vertex_global_url_support.py b/tests/unit/llms/vertex_ai/test_vertex_global_url_support.py index 5a007f5a4f6..76d8ccb3ecd 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_global_url_support.py +++ b/tests/unit/llms/vertex_ai/test_vertex_global_url_support.py @@ -15,8 +15,8 @@ import pytest from litellm.llms.vertex_ai.common_utils import ( _get_embedding_url, - _get_vertex_url, get_vertex_base_url, + get_vertex_url, ) @@ -80,7 +80,7 @@ class TestChatCompletionURLs: "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, endpoint = _get_vertex_url( + url, endpoint = get_vertex_url( mode="chat", model="gemini-1.5-pro", stream=stream, @@ -112,7 +112,7 @@ class TestChatCompletionURLs: "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, endpoint = _get_vertex_url( + url, endpoint = get_vertex_url( mode="chat", model="1234567890", # Numeric model ID stream=stream, @@ -229,7 +229,7 @@ class TestCountTokensURLs: "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, endpoint = _get_vertex_url( + url, endpoint = get_vertex_url( mode="count_tokens", model="gemini-1.5-pro", stream=None, @@ -274,15 +274,13 @@ class TestImageGenerationURLs: ), ], ) - def test_image_generation_url_construction( - self, vertex_location, model, expected_url_pattern - ): + def test_image_generation_url_construction(self, vertex_location, model, expected_url_pattern): """Test that image_generation URLs are correctly constructed for regional and global locations.""" with patch( "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, endpoint = _get_vertex_url( + url, endpoint = get_vertex_url( mode="image_generation", model=model, stream=None, @@ -313,7 +311,7 @@ class TestAPIVersions: "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, _ = _get_vertex_url( + url, _ = get_vertex_url( mode="chat", model="gemini-1.5-pro", stream=False, @@ -354,7 +352,7 @@ class TestEdgeCases: vertex_api_version="v1", ) else: - url, _ = _get_vertex_url( + url, _ = get_vertex_url( mode=mode, model="gemini-1.5-pro", stream=False, @@ -376,7 +374,7 @@ class TestEdgeCases: "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, _ = _get_vertex_url( + url, _ = get_vertex_url( mode="chat", model="gemini-1.5-pro", stream=False, @@ -398,7 +396,7 @@ class TestBackwardCompatibility: "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, _ = _get_vertex_url( + url, _ = get_vertex_url( mode="chat", model="gemini-1.5-pro", stream=False, @@ -419,7 +417,7 @@ class TestBackwardCompatibility: "litellm.VertexGeminiConfig.get_model_for_vertex_ai_url", side_effect=lambda model: model, ): - url, _ = _get_vertex_url( + url, _ = get_vertex_url( mode="chat", model="gemini-1.5-pro", stream=True, diff --git a/tests/unit/llms/vertex_ai/test_vertex_llm_base.py b/tests/unit/llms/vertex_ai/test_vertex_llm_base.py index a4d67606698..b5a24ab02d9 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/unit/llms/vertex_ai/test_vertex_llm_base.py @@ -34,17 +34,15 @@ class TestVertexBase: mock_creds.quota_project_id = "project-1" # Test case 1: Ensure credentials match project - with patch.object( - vertex_base, "load_auth", return_value=(mock_creds, "project-1") - ): + with patch.object(vertex_base, "load_auth", return_value=(mock_creds, "project-1")): if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials={"type": "service_account", "project_id": "project-1"}, project_id="project-1", custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials={"type": "service_account", "project_id": "project-1"}, project_id="project-1", custom_llm_provider="vertex_ai", @@ -53,17 +51,15 @@ class TestVertexBase: assert token == "fake-token-1" # Test case 2: Allow using credentials from different project - with patch.object( - vertex_base, "load_auth", return_value=(mock_creds, "project-1") - ): + with patch.object(vertex_base, "load_auth", return_value=(mock_creds, "project-1")): if is_async: - result = await vertex_base._ensure_access_token_async( + result = await vertex_base.ensure_access_token_async( credentials={"type": "service_account"}, project_id="different-project", custom_llm_provider="vertex_ai", ) else: - result = vertex_base._ensure_access_token( + result = vertex_base.ensure_access_token( credentials={"type": "service_account"}, project_id="different-project", custom_llm_provider="vertex_ai", @@ -83,18 +79,16 @@ class TestVertexBase: mock_creds.quota_project_id = "project-1" # Test initial credential load and caching - with patch.object( - vertex_base, "load_auth", return_value=(mock_creds, "project-1") - ): + with patch.object(vertex_base, "load_auth", return_value=(mock_creds, "project-1")): # First call should load credentials if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials={"type": "service_account"}, project_id="project-1", custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials={"type": "service_account"}, project_id="project-1", custom_llm_provider="vertex_ai", @@ -103,13 +97,13 @@ class TestVertexBase: # Second call should use cached credentials if is_async: - token2, project2 = await vertex_base._ensure_access_token_async( + token2, project2 = await vertex_base.ensure_access_token_async( credentials={"type": "service_account"}, project_id="project-1", custom_llm_provider="vertex_ai", ) else: - token2, project2 = vertex_base._ensure_access_token( + token2, project2 = vertex_base.ensure_access_token( credentials={"type": "service_account"}, project_id="project-1", custom_llm_provider="vertex_ai", @@ -143,13 +137,13 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials={"type": "service_account"}, project_id="project-1", custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials={"type": "service_account"}, project_id="project-1", custom_llm_provider="vertex_ai", @@ -166,11 +160,11 @@ class TestVertexBase: # Test that Gemini requests bypass credential checks if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=None, project_id=None, custom_llm_provider="gemini" ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials=None, project_id=None, custom_llm_provider="gemini" ) assert token == "" @@ -214,13 +208,13 @@ class TestVertexBase: # 1. Test that authorized_user-style credentials are correctly handled and uses quota_project_id if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", @@ -233,13 +227,13 @@ class TestVertexBase: # 2. Test that authorized_user-style credentials are correctly handled and uses passed in project_id not_quota_project_id = "new-project" if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=not_quota_project_id, custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials=credentials, project_id=not_quota_project_id, custom_llm_provider="vertex_ai", @@ -278,13 +272,13 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl if is_async: - token, _ = await vertex_base._ensure_access_token_async( + token, _ = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) else: - token, _ = vertex_base._ensure_access_token( + token, _ = vertex_base.ensure_access_token( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", @@ -329,22 +323,22 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl if is_async: - token, _ = await vertex_base._ensure_access_token_async( + token, _ = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) else: - token, _ = vertex_base._ensure_access_token( + token, _ = vertex_base.ensure_access_token( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) assert mock_credentials_from_identity_pool_with_aws.called - assert mock_credentials_from_identity_pool_with_aws.call_args[1][ - "scopes" - ] == ["https://www.googleapis.com/auth/cloud-platform"] + assert mock_credentials_from_identity_pool_with_aws.call_args[1]["scopes"] == [ + "https://www.googleapis.com/auth/cloud-platform" + ] assert token == "refreshed-token" @pytest.mark.parametrize("is_async", [True, False], ids=["async", "sync"]) @@ -361,17 +355,15 @@ class TestVertexBase: credentials = {"type": "service_account", "project_id": "project-1"} - with patch.object( - vertex_base, "load_auth", return_value=(mock_creds, "project-1") - ): + with patch.object(vertex_base, "load_auth", return_value=(mock_creds, "project-1")): if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -410,13 +402,13 @@ class TestVertexBase: # Should handle old format gracefully if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -439,18 +431,16 @@ class TestVertexBase: credentials = {"type": "service_account"} - with patch.object( - vertex_base, "load_auth", return_value=(mock_creds, "resolved-project") - ): + with patch.object(vertex_base, "load_auth", return_value=(mock_creds, "resolved-project")): # Call without project_id, should use resolved project from credentials if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", @@ -509,13 +499,13 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl if is_async: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", ) else: - token, project = vertex_base._ensure_access_token( + token, project = vertex_base.ensure_access_token( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -549,18 +539,16 @@ class TestVertexBase: credentials = {"type": "service_account", "project_id": "cred-project"} - with patch.object( - vertex_base, "load_auth", return_value=(mock_creds, "cred-project") - ): + with patch.object(vertex_base, "load_auth", return_value=(mock_creds, "cred-project")): # First call with explicit project_id if is_async: - token1, project1 = await vertex_base._ensure_access_token_async( + token1, project1 = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="explicit-project", custom_llm_provider="vertex_ai", ) else: - token1, project1 = vertex_base._ensure_access_token( + token1, project1 = vertex_base.ensure_access_token( credentials=credentials, project_id="explicit-project", custom_llm_provider="vertex_ai", @@ -568,13 +556,13 @@ class TestVertexBase: # Second call with None project_id (should use credential project) if is_async: - token2, project2 = await vertex_base._ensure_access_token_async( + token2, project2 = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) else: - token2, project2 = vertex_base._ensure_access_token( + token2, project2 = vertex_base.ensure_access_token( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", @@ -617,16 +605,15 @@ class TestVertexBase: "load_auth", return_value=(mock_creds, "resolved-from-credentials"), ) as mock_load_auth: - # First call: User provides NO project_id, should resolve from credentials if is_async: - token1, project1 = await vertex_base._ensure_access_token_async( + token1, project1 = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, # Key: user doesn't provide project_id custom_llm_provider="vertex_ai", ) else: - token1, project1 = vertex_base._ensure_access_token( + token1, project1 = vertex_base.ensure_access_token( credentials=credentials, project_id=None, # Key: user doesn't provide project_id custom_llm_provider="vertex_ai", @@ -659,13 +646,13 @@ class TestVertexBase: # Second call: Same scenario - should use cache and NOT call load_auth again if is_async: - token2, project2 = await vertex_base._ensure_access_token_async( + token2, project2 = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, # Still no project_id provided custom_llm_provider="vertex_ai", ) else: - token2, project2 = vertex_base._ensure_access_token( + token2, project2 = vertex_base.ensure_access_token( credentials=credentials, project_id=None, # Still no project_id provided custom_llm_provider="vertex_ai", @@ -679,13 +666,13 @@ class TestVertexBase: # Third call: Now user provides the resolved project_id explicitly # This should also use cache (the resolved_cache_key) if is_async: - token3, project3 = await vertex_base._ensure_access_token_async( + token3, project3 = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="resolved-from-credentials", # Explicit resolved project_id custom_llm_provider="vertex_ai", ) else: - token3, project3 = vertex_base._ensure_access_token( + token3, project3 = vertex_base.ensure_access_token( credentials=credentials, project_id="resolved-from-credentials", # Explicit resolved project_id custom_llm_provider="vertex_ai", @@ -1582,13 +1569,13 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl if is_async: - token, _ = await vertex_base._ensure_access_token_async( + token, _ = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) else: - token, _ = vertex_base._ensure_access_token( + token, _ = vertex_base.ensure_access_token( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", @@ -1645,13 +1632,13 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl if is_async: - token, _ = await vertex_base._ensure_access_token_async( + token, _ = await vertex_base.ensure_access_token_async( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", ) else: - token, _ = vertex_base._ensure_access_token( + token, _ = vertex_base.ensure_access_token( credentials=credentials, project_id=None, custom_llm_provider="vertex_ai", @@ -1767,7 +1754,7 @@ class TestVertexBase: # Launch 50 concurrent requests tasks = [ - vertex_base._ensure_access_token_async( + vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -1829,7 +1816,7 @@ class TestVertexBase: ): results = await asyncio.gather( *[ - vertex_base._ensure_access_token_async( + vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -1875,7 +1862,7 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -1917,7 +1904,7 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -1946,7 +1933,7 @@ class TestVertexBase: ) with patch.object(vertex_base, "refresh_auth") as mock_refresh: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -1986,7 +1973,7 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl - await vertex_base._ensure_access_token_async( + await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", @@ -2037,7 +2024,7 @@ class TestVertexBase: mock_refresh.side_effect = mock_refresh_impl - await vertex_base._ensure_access_token_async( + await vertex_base.ensure_access_token_async( credentials=credentials, project_id=f"project-{i}", custom_llm_provider="vertex_ai", @@ -2098,7 +2085,7 @@ class TestVertexBase: ), patch.object(vertex_base, "refresh_auth"), ): - await vertex_base._ensure_access_token_async( + await vertex_base.ensure_access_token_async( credentials=credentials, project_id=f"project-{i}", custom_llm_provider="vertex_ai", @@ -2185,7 +2172,7 @@ class TestVertexBase: "_acquire_async_refresh_lock", wraps=vertex_base._acquire_async_refresh_lock, ) as mock_get_lock: - token, project = await vertex_base._ensure_access_token_async( + token, project = await vertex_base.ensure_access_token_async( credentials=credentials, project_id="project-1", custom_llm_provider="vertex_ai", diff --git a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py index efe97ce33a8..1110b53624a 100644 --- a/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py +++ b/tests/unit/llms/vertex_ai/vertex_gemma_models/test_vertex_gemma_transformation.py @@ -413,7 +413,7 @@ class TestVertexGemmaCompletion: } with ( - patch("litellm.llms.vertex_ai.vertex_gemma_models.transformation._get_httpx_client") as mock_get_client, + patch("litellm.llms.vertex_ai.vertex_gemma_models.transformation.get_httpx_client") as mock_get_client, patch( "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", return_value=("fake-access-token", "PROJECT_ID"), @@ -929,7 +929,7 @@ class TestVertexGemmaCompletion: ) with ( - patch("litellm.llms.vertex_ai.vertex_gemma_models.transformation._get_httpx_client") as mock_get_client, + patch("litellm.llms.vertex_ai.vertex_gemma_models.transformation.get_httpx_client") as mock_get_client, patch( "litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token", return_value=("fake-access-token", "PROJECT_ID"), diff --git a/tests/unit/llms/watsonx/test_watsonx.py b/tests/unit/llms/watsonx/test_watsonx.py index dc75bdff6b1..14a905dcd23 100644 --- a/tests/unit/llms/watsonx/test_watsonx.py +++ b/tests/unit/llms/watsonx/test_watsonx.py @@ -41,8 +41,10 @@ async def test_watsonx_text_gpt_oss_async_completion_fetches_hf_template_off_the return httpx.Response(200, content=chat_template.encode()) return httpx.Response(200, json={"chat_template": chat_template, "bos_token": None, "eos_token": None}) - monkeypatch.setattr(huggingface_template_handler, "_get_httpx_client", forbid_sync_client) - monkeypatch.setattr(huggingface_template_handler, "get_async_httpx_client", lambda **kwargs: Mock(get=serve_hf_file)) + monkeypatch.setattr(huggingface_template_handler, "get_httpx_client", forbid_sync_client) + monkeypatch.setattr( + huggingface_template_handler, "get_async_httpx_client", lambda **kwargs: Mock(get=serve_hf_file) + ) def handle(request): captured["body"] = json.loads(request.content) diff --git a/tests/unit/llms/xai/test_xai_chat_transformation.py b/tests/unit/llms/xai/test_xai_chat_transformation.py index 601b384aae4..f77f6e243b1 100644 --- a/tests/unit/llms/xai/test_xai_chat_transformation.py +++ b/tests/unit/llms/xai/test_xai_chat_transformation.py @@ -153,14 +153,14 @@ class TestXAIUsageNormalization: def test_preserves_reasoning_tokens_in_total_usage(self): usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=200) - XAIChatConfig._normalize_openai_compatible_usage_totals(usage) + XAIChatConfig.normalize_openai_compatible_usage_totals(usage) assert usage.total_tokens == 200 def test_preserves_reasoning_tokens_in_streaming_usage(self): usage = {"prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 200} - XAIChatConfig._normalize_openai_compatible_usage_totals(usage) + XAIChatConfig.normalize_openai_compatible_usage_totals(usage) assert usage["total_tokens"] == 200 diff --git a/tests/unit/llms/xai/test_xai_key_fallback.py b/tests/unit/llms/xai/test_xai_key_fallback.py index cbc507c5ee3..c7f3bd39920 100644 --- a/tests/unit/llms/xai/test_xai_key_fallback.py +++ b/tests/unit/llms/xai/test_xai_key_fallback.py @@ -90,7 +90,7 @@ def test_chat_config_uses_xai_key_fallback(monkeypatch): monkeypatch.setattr(litellm, "api_key", None) monkeypatch.delenv("XAI_API_KEY", raising=False) - _, api_key = XAIChatConfig()._get_openai_compatible_provider_info(None, None) + _, api_key = XAIChatConfig().get_openai_compatible_provider_info(None, None) assert api_key == "xai_key_value" @@ -100,7 +100,7 @@ def test_chat_config_uses_environment_key_fallback(monkeypatch): monkeypatch.setattr(litellm, "api_key", None) monkeypatch.setenv("XAI_API_KEY", "env_api_key") - _, api_key = XAIChatConfig()._get_openai_compatible_provider_info(None, None) + _, api_key = XAIChatConfig().get_openai_compatible_provider_info(None, None) assert api_key == "env_api_key" @@ -110,7 +110,7 @@ def test_chat_config_does_not_use_generic_key_fallback(monkeypatch): monkeypatch.setattr(litellm, "api_key", "common_api_key") monkeypatch.delenv("XAI_API_KEY", raising=False) - _, api_key = XAIChatConfig()._get_openai_compatible_provider_info(None, None) + _, api_key = XAIChatConfig().get_openai_compatible_provider_info(None, None) assert api_key is None @@ -120,9 +120,7 @@ def test_chat_config_prefers_explicit_api_key(monkeypatch): monkeypatch.setattr(litellm, "api_key", "common_api_key") monkeypatch.setenv("XAI_API_KEY", "env_api_key") - _, api_key = XAIChatConfig()._get_openai_compatible_provider_info( - None, "param_api_key" - ) + _, api_key = XAIChatConfig().get_openai_compatible_provider_info(None, "param_api_key") assert api_key == "param_api_key" diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 79b0828ec15..14fd60c9793 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -588,7 +588,7 @@ class TestVertexAIPassThroughHandler: with ( mock.patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.vertex_llm_base._ensure_access_token_async" + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.vertex_llm_base.ensure_access_token_async" ) as mock_ensure_token, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.vertex_llm_base._get_token_and_url" @@ -5502,7 +5502,7 @@ class TestVertexAILiveWebsocketPassthrough: ws_passthrough = AsyncMock() with ( - patch.object(passthrough_module.vertex_llm_base, "_ensure_access_token_async", ensure_token), + patch.object(passthrough_module.vertex_llm_base, "ensure_access_token_async", ensure_token), patch.object(passthrough_module, "websocket_passthrough_request", ws_passthrough), ): await passthrough_module.vertex_ai_live_websocket_passthrough( @@ -5611,7 +5611,7 @@ class TestVertexAILiveWebsocketPassthrough: ws_passthrough = AsyncMock() with ( - patch.object(passthrough_module.vertex_llm_base, "_ensure_access_token_async", ensure_token), + patch.object(passthrough_module.vertex_llm_base, "ensure_access_token_async", ensure_token), patch.object(passthrough_module, "websocket_passthrough_request", ws_passthrough), ): await passthrough_module.vertex_ai_live_websocket_passthrough( @@ -5643,7 +5643,7 @@ class TestVertexAILiveWebsocketPassthrough: ensure_token = AsyncMock(side_effect=Exception("Unable to find your credentials")) with ( - patch.object(passthrough_module.vertex_llm_base, "_ensure_access_token_async", ensure_token), + patch.object(passthrough_module.vertex_llm_base, "ensure_access_token_async", ensure_token), patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging, ): mock_proxy_logging.post_call_failure_hook = AsyncMock() diff --git a/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py b/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py index 46f024366ec..0a84ee55a89 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py +++ b/tests/unit/proxy/pass_through_endpoints/test_vertex_ai_batch_passthrough.py @@ -86,9 +86,7 @@ class TestVertexAIBatchPassthroughHandler: "error_file_id": None, "completion_window": "24h", } - mock_transformation._get_batch_id_from_vertex_ai_batch_response.return_value = ( - "123456789" - ) + mock_transformation.get_batch_id_from_vertex_ai_batch_response.return_value = "123456789" # Test the handler result = VertexPassthroughLoggingHandler.batch_prediction_jobs_handler( @@ -445,7 +443,7 @@ class TestVertexAIBatchPassthroughHandler: "input_file_id": "gs://bucket/in.jsonl", "completion_window": "24h", } - mock_transformation._get_batch_id_from_vertex_ai_batch_response.return_value = "123456" + mock_transformation.get_batch_id_from_vertex_ai_batch_response.return_value = "123456" VertexPassthroughLoggingHandler.batch_prediction_jobs_handler( httpx_response=response, @@ -507,11 +505,7 @@ class TestVertexAIBatchPassthroughHandler: expected_results = ["456789", "def123", "999", "invalid-format"] for test_case, expected in zip(test_cases, expected_results): - result = ( - VertexAIBatchTransformation._get_batch_id_from_vertex_ai_batch_response( - {"name": test_case} - ) - ) + result = VertexAIBatchTransformation.get_batch_id_from_vertex_ai_batch_response({"name": test_case}) assert result == expected def test_model_name_extraction_from_vertex_path(self): @@ -576,9 +570,7 @@ class TestVertexAIBatchPassthroughHandler: "error_file_id": None, "completion_window": "24h", } - mock_transformation._get_batch_id_from_vertex_ai_batch_response.return_value = ( - "123456789" - ) + mock_transformation.get_batch_id_from_vertex_ai_batch_response.return_value = "123456789" # Test the complete workflow result = VertexPassthroughLoggingHandler.batch_prediction_jobs_handler( diff --git a/tests/unit/proxy/public_endpoints/test_public_endpoints.py b/tests/unit/proxy/public_endpoints/test_public_endpoints.py index 448ad9c712e..99a76d59448 100644 --- a/tests/unit/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/unit/proxy/public_endpoints/test_public_endpoints.py @@ -499,7 +499,7 @@ def test_google_ai_studio_provider_fields_expose_api_base(): assert api_base_field["field_type"] == "text" # default_value MUST be null (not the canonical URL): saving it as the # default would persist v1beta into every credential record and bypass - # `_get_gemini_url`'s automatic v1alpha routing for Gemini 3+ models. The + # `get_gemini_url`'s automatic v1alpha routing for Gemini 3+ models. The # placeholder shows the canonical URL so users still get the visual hint. # (See greptileai threads on PR #30419.) assert api_base_field["default_value"] is None diff --git a/tests/unit/router_utils/test_reasoning_effort_capability.py b/tests/unit/router_utils/test_reasoning_effort_capability.py index a08c99e269c..6a0c274f149 100644 --- a/tests/unit/router_utils/test_reasoning_effort_capability.py +++ b/tests/unit/router_utils/test_reasoning_effort_capability.py @@ -200,7 +200,7 @@ class TestNoneLevelPolarity: resolved = resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) assert resolved is not None - gate_accepts_none = AzureOpenAIGPT5Config._supports_reasoning_effort_level(model_key, "none") + gate_accepts_none = AzureOpenAIGPT5Config.supports_reasoning_effort_level(model_key, "none") assert ("none" in resolved) is gate_accepts_none diff --git a/tests/unit/test_aembedding_session_reuse_e2e.py b/tests/unit/test_aembedding_session_reuse_e2e.py index 15662d4d35c..c5605bc6dff 100644 --- a/tests/unit/test_aembedding_session_reuse_e2e.py +++ b/tests/unit/test_aembedding_session_reuse_e2e.py @@ -26,11 +26,11 @@ def test_openai_embedding_passes_shared_session(): Verify shared_session flows through the complete call chain. Full chain: litellm.embedding() -> OpenAI.embedding() -> _get_openai_client() - -> AsyncHTTPHandler -> _create_async_transport() -> _create_aiohttp_transport() + -> AsyncHTTPHandler -> create_async_transport() -> create_aiohttp_transport() """ import litellm - from litellm.llms.openai.openai import OpenAIChatCompletion from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.llms.openai.openai import OpenAIChatCompletion # Step 1: litellm.embedding() extracts and passes shared_session main_source = inspect.getsource(litellm.embedding) @@ -46,16 +46,14 @@ def test_openai_embedding_passes_shared_session(): client_source = inspect.getsource(OpenAIChatCompletion._get_openai_client) assert "shared_session" in client_source - # Step 4: AsyncHTTPHandler.create_client passes it to _create_async_transport + # Step 4: AsyncHTTPHandler.create_client passes it to create_async_transport create_client_source = inspect.getsource(AsyncHTTPHandler.create_client) assert "shared_session=shared_session" in create_client_source - # Step 5: _create_async_transport passes it to _create_aiohttp_transport - async_transport_source = inspect.getsource(AsyncHTTPHandler._create_async_transport) + # Step 5: create_async_transport passes it to create_aiohttp_transport + async_transport_source = inspect.getsource(AsyncHTTPHandler.create_async_transport) assert "shared_session=shared_session" in async_transport_source - # Step 6: _create_aiohttp_transport uses it - aiohttp_transport_source = inspect.getsource( - AsyncHTTPHandler._create_aiohttp_transport - ) + # Step 6: create_aiohttp_transport uses it + aiohttp_transport_source = inspect.getsource(AsyncHTTPHandler.create_aiohttp_transport) assert "shared_session" in aiohttp_transport_source diff --git a/tests/unit/test_claude_fable_5_config.py b/tests/unit/test_claude_fable_5_config.py index dfbda795c7a..7ec76742bb8 100644 --- a/tests/unit/test_claude_fable_5_config.py +++ b/tests/unit/test_claude_fable_5_config.py @@ -48,7 +48,7 @@ def test_adaptive_thinking_detected_for_fable_5(local_model_cost_map, model): maps to ``thinking.type='adaptive'`` + ``output_config.effort``.""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is True FABLE_5_1_VARIANTS = ( @@ -91,6 +91,4 @@ def test_fable_5_1_registered_for_bedrock_converse(): def test_adaptive_thinking_detected_for_fable_5_1(local_model_cost_map, model): from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, "anthropic") is True - - + assert AnthropicModelInfo.is_adaptive_thinking_model(model, "anthropic") is True diff --git a/tests/unit/test_claude_opus_5_config.py b/tests/unit/test_claude_opus_5_config.py index b5efe016a60..bac4beb494c 100644 --- a/tests/unit/test_claude_opus_5_config.py +++ b/tests/unit/test_claude_opus_5_config.py @@ -103,6 +103,6 @@ def test_opus_5_5_thinking_profile(local_model_cost_map, model, provider): no forced tool use, same as Fable 5.1.""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, provider) is True + assert AnthropicModelInfo.is_adaptive_thinking_model(model, provider) is True assert AnthropicModelInfo._is_always_on_thinking_model(model, provider) is True assert AnthropicModelInfo.forced_tool_use_unsupported(model.removeprefix("anthropic/")) is True diff --git a/tests/unit/test_claude_sonnet_5_config.py b/tests/unit/test_claude_sonnet_5_config.py index 36078ff494f..89ad89d8912 100644 --- a/tests/unit/test_claude_sonnet_5_config.py +++ b/tests/unit/test_claude_sonnet_5_config.py @@ -82,6 +82,6 @@ def test_sonnet_5_5_thinking_profile(local_model_cost_map, model, provider): no forced tool use, same as Opus 5.5.""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo - assert AnthropicModelInfo._is_adaptive_thinking_model(model, provider) is True + assert AnthropicModelInfo.is_adaptive_thinking_model(model, provider) is True assert AnthropicModelInfo._is_always_on_thinking_model(model, provider) is True assert AnthropicModelInfo.forced_tool_use_unsupported(model.removeprefix("anthropic/")) is True diff --git a/tests/unit/test_constants.py b/tests/unit/test_constants.py index e16ea6617d4..18ed6975605 100644 --- a/tests/unit/test_constants.py +++ b/tests/unit/test_constants.py @@ -193,17 +193,17 @@ def test_morph_config_get_provider_info(): # Test with environment variable with patch.dict(os.environ, {"MORPH_API_KEY": "test-key-from-env"}): - api_base, api_key = config._get_openai_compatible_provider_info(None, None) + api_base, api_key = config.get_openai_compatible_provider_info(None, None) assert api_base == "https://api.morphllm.com/v1" assert api_key == "test-key-from-env" # Test with passed api_key - api_base, api_key = config._get_openai_compatible_provider_info(None, "direct-key") + api_base, api_key = config.get_openai_compatible_provider_info(None, "direct-key") assert api_base == "https://api.morphllm.com/v1" assert api_key == "direct-key" # Test with custom api_base - api_base, api_key = config._get_openai_compatible_provider_info("https://custom.morph.com", "key") + api_base, api_key = config.get_openai_compatible_provider_info("https://custom.morph.com", "key") assert api_base == "https://custom.morph.com" assert api_key == "key" diff --git a/tests/unit/test_cost_calculator.py b/tests/unit/test_cost_calculator.py index 9e348c0579a..cc96edd78b1 100644 --- a/tests/unit/test_cost_calculator.py +++ b/tests/unit/test_cost_calculator.py @@ -2686,7 +2686,7 @@ def test_gemini_cache_tokens_details_no_negative_values(): } } - usage = VertexGeminiConfig._calculate_usage(completion_response) + usage = VertexGeminiConfig.calculate_usage(completion_response) # Text tokens should be non-cached text only: 9402 - 9393 = 9 assert usage.prompt_tokens_details.text_tokens == 9, ( @@ -2735,7 +2735,7 @@ def test_gemini_without_cache_tokens_details(): } } - usage = VertexGeminiConfig._calculate_usage(completion_response) + usage = VertexGeminiConfig.calculate_usage(completion_response) # Should use promptTokensDetails values directly assert usage.prompt_tokens_details.text_tokens == 6 diff --git a/tests/unit/test_private_usage_aliases.py b/tests/unit/test_private_usage_aliases.py index a3fa4ac27c3..412cae4eb0c 100644 --- a/tests/unit/test_private_usage_aliases.py +++ b/tests/unit/test_private_usage_aliases.py @@ -1,6 +1,8 @@ +import inspect +from collections.abc import Awaitable, Callable from importlib import import_module from types import MethodType -from typing import Final +from typing import Final, cast import pytest @@ -972,10 +974,1261 @@ PROPERTY_CASES: Final = ( ) CLASS_PROPERTY_CASES: Final = () +LLMS_ALIAS_CASES: Final = ( + ( + "litellm.llms.anthropic.pass_through.adapters.transformation", + "LiteLLMAnthropicMessagesAdapter", + "_translate_openai_usage_to_anthropic_usage", + "translate_openai_usage_to_anthropic_usage", + False, + ), + ( + "litellm.llms.anthropic.pass_through.adapters.transformation", + "LiteLLMAnthropicMessagesAdapter", + "_translate_openai_usage_to_anthropic_usage_delta", + "translate_openai_usage_to_anthropic_usage_delta", + False, + ), + ( + "litellm.llms.anthropic.pass_through.adapters.transformation", + "LiteLLMAnthropicMessagesAdapter", + "_translate_streaming_openai_chunk_to_anthropic_content_block", + "translate_streaming_openai_chunk_to_anthropic_content_block", + False, + ), + ( + "litellm.llms.anthropic.pass_through.context_management.dispatcher", + "", + "_normalize_spec", + "normalize_spec", + False, + ), + ( + "litellm.llms.anthropic.pass_through.messages.streaming_iterator", + "", + "_is_message_stop_chunk", + "is_message_stop_chunk", + False, + ), + ( + "litellm.llms.anthropic.pass_through.messages.streaming_iterator", + "", + "_is_provider_error_chunk", + "is_provider_error_chunk", + False, + ), + ( + "litellm.llms.anthropic.pass_through.messages.transformation", + "", + "_messages_carry_output_config", + "messages_carry_output_config", + False, + ), + ("litellm.llms.azure.azure", "", "_check_dynamic_azure_params", "check_dynamic_azure_params", False), + ( + "litellm.llms.base_llm.base_utils", + "", + "_convert_tool_response_to_message", + "convert_tool_response_to_message", + False, + ), + ( + "litellm.llms.bedrock.batches.handler", + "BedrockBatchesHandler", + "_handle_async_invoke_status", + "handle_async_invoke_status", + False, + ), + ( + "litellm.llms.bedrock.batches.handler", + "BedrockBatchesHandler", + "_handle_model_invocation_job_status", + "handle_model_invocation_job_status", + False, + ), + ("litellm.llms.bedrock.common_utils", "", "_get_all_bedrock_regions", "get_all_bedrock_regions", False), + ( + "litellm.llms.custom_httpx.aiohttp_transport", + "LiteLLMAiohttpTransport", + "_get_valid_client_session", + "get_valid_client_session", + False, + ), + ( + "litellm.llms.custom_httpx.http_handler", + "", + "_build_aiohttp_keepalive_socket_factory", + "build_aiohttp_keepalive_socket_factory", + False, + ), + ("litellm.llms.custom_httpx.http_handler", "", "_get_httpx_client", "get_httpx_client", False), + ( + "litellm.llms.custom_httpx.http_handler", + "AsyncHTTPHandler", + "_create_aiohttp_transport", + "create_aiohttp_transport", + False, + ), + ( + "litellm.llms.custom_httpx.http_handler", + "AsyncHTTPHandler", + "_create_async_transport", + "create_async_transport", + False, + ), + ( + "litellm.llms.custom_httpx.http_handler", + "AsyncHTTPHandler", + "_create_httpx_proxy_mounts", + "create_httpx_proxy_mounts", + False, + ), + ( + "litellm.llms.custom_httpx.http_handler", + "AsyncHTTPHandler", + "_should_use_aiohttp_transport", + "should_use_aiohttp_transport", + False, + ), + ( + "litellm.llms.huggingface.common_utils", + "", + "_fetch_inference_provider_mapping", + "fetch_inference_provider_mapping", + False, + ), + ("litellm.llms.oci.chat.cohere", "", "_extract_text_content", "extract_text_content", False), + ("litellm.llms.oci.chat.generic", "", "_normalize_oci_finish_reason", "normalize_oci_finish_reason", False), + ("litellm.llms.oci.chat.generic", "", "_synthesize_oci_tool_call_id", "synthesize_oci_tool_call_id", False), + ("litellm.llms.ollama.common_utils", "", "_convert_image", "convert_image", False), + ("litellm.llms.openai.chat.gpt_5_transformation", "", "_get_effort_level", "get_effort_level", False), + ("litellm.llms.openai.completion.utils", "", "_transform_prompt", "transform_prompt", False), + ( + "litellm.llms.openai.cost_calculation", + "", + "_video_output_cost_per_second", + "video_output_cost_per_second", + False, + ), + ( + "litellm.llms.openai.fine_tuning.handler", + "", + "_litellm_fine_tuning_job_from_response", + "litellm_fine_tuning_job_from_response", + False, + ), + ( + "litellm.llms.vertex_ai.batches.transformation", + "VertexAIBatchTransformation", + "_get_batch_id_from_vertex_ai_batch_response", + "get_batch_id_from_vertex_ai_batch_response", + False, + ), + ("litellm.llms.vertex_ai.common_utils", "", "_build_json_schema", "build_json_schema", False), + ("litellm.llms.vertex_ai.common_utils", "", "_build_vertex_schema", "build_vertex_schema", False), + ("litellm.llms.vertex_ai.common_utils", "", "_check_text_in_content", "check_text_in_content", False), + ( + "litellm.llms.vertex_ai.common_utils", + "", + "_convert_vertex_datetime_to_openai_datetime", + "convert_vertex_datetime_to_openai_datetime", + False, + ), + ("litellm.llms.vertex_ai.common_utils", "", "_get_gemini_url", "get_gemini_url", False), + ("litellm.llms.vertex_ai.common_utils", "", "_get_vertex_url", "get_vertex_url", False), + ("litellm.llms.vertex_ai.gemini.transformation", "", "_camel_to_snake", "camel_to_snake", False), + ( + "litellm.llms.vertex_ai.gemini.transformation", + "", + "_gemini_convert_messages_with_history", + "gemini_convert_messages_with_history", + False, + ), + ("litellm.llms.vertex_ai.gemini.transformation", "", "_snake_to_camel", "snake_to_camel", False), + ( + "litellm.llms.vertex_ai.gemini.transformation", + "", + "_transform_request_body", + "transform_request_body", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.transformation", + "", + "_transform_system_message", + "transform_system_message", + False, + ), + ( + "litellm.llms.vertex_ai.gemini_embeddings.batch_embed_content_transformation", + "", + "_is_file_reference", + "is_file_reference", + False, + ), + ( + "litellm.llms.vertex_ai.vertex_llm_base", + "", + "_graft_default_vertex_path", + "graft_default_vertex_path", + False, + ), + ("litellm.llms.watsonx.common_utils", "", "_generate_watsonx_token", "generate_watsonx_token", False), + ("litellm.llms.watsonx.common_utils", "", "_get_api_params", "get_api_params", False), +) +LLMS_FORWARDER_CASES: Final = ( + ( + "litellm.llms.bedrock.chat.invoke_handler", + "AWSEventStreamDecoder", + "_chunk_parser", + "chunk_parser", + "instance", + False, + ), + ( + "litellm.llms.aiml.chat.transformation", + "AIMLChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_convert_tool_response_to_message", + "convert_tool_response_to_message", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_map_reasoning_effort", + "map_reasoning_effort", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_map_stop_sequences", + "map_stop_sequences", + "instance", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_map_tool_helper", + "map_tool_helper", + "instance", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_map_tools", + "map_tools", + "instance", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_maybe_drop_speed_param", + "maybe_drop_speed_param", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_model_supports_effort_param", + "model_supports_effort_param", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_raise_invalid_reasoning_effort", + "raise_invalid_reasoning_effort", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.chat.transformation", + "AnthropicConfig", + "_validate_effort_for_model", + "validate_effort_for_model", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.common_utils", + "AnthropicModelInfo", + "_apply_sampling_param", + "apply_sampling_param", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.common_utils", + "AnthropicModelInfo", + "_is_adaptive_thinking_model", + "is_adaptive_thinking_model", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.common_utils", + "AnthropicModelInfo", + "_supports_model_capability", + "supports_model_capability", + "staticmethod", + False, + ), + ( + "litellm.llms.anthropic.completion.transformation", + "AnthropicTextConfig", + "_is_anthropic_text_model", + "is_anthropic_text_model", + "staticmethod", + False, + ), + ( + "litellm.llms.azure.common_utils", + "BaseAzureLLM", + "_base_validate_azure_environment", + "base_validate_azure_environment", + "staticmethod", + False, + ), + ( + "litellm.llms.azure.common_utils", + "BaseAzureLLM", + "_get_base_azure_url", + "get_base_azure_url", + "staticmethod", + False, + ), + ( + "litellm.llms.azure.common_utils", + "BaseAzureLLM", + "_try_get_default_azure_credential_provider", + "try_get_default_azure_credential_provider", + "staticmethod", + False, + ), + ( + "litellm.llms.azure.realtime.handler", + "AzureOpenAIRealtime", + "_construct_url", + "construct_url", + "instance", + False, + ), + ( + "litellm.llms.azure_ai.agents.transformation", + "AzureAIAgentsConfig", + "_get_agent_id", + "get_agent_id", + "instance", + False, + ), + ( + "litellm.llms.azure_ai.chat.transformation", + "AzureAIStudioConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.azure_ai.embed.cohere_transformation", + "AzureAICohereConfig", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.azure_ai.embed.cohere_transformation", + "AzureAICohereConfig", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.base_llm.base_model_iterator", + "BaseModelResponseIterator", + "_string_to_dict_parser", + "string_to_dict_parser", + "staticmethod", + False, + ), + ( + "litellm.llms.base_llm.chat.transformation", + "BaseConfig", + "_add_tools_to_optional_params", + "add_tools_to_optional_params", + "instance", + False, + ), + ( + "litellm.llms.base_llm.passthrough.transformation", + "BasePassthroughConfig", + "_convert_raw_bytes_to_str_lines", + "convert_raw_bytes_to_str_lines", + "instance", + False, + ), + ( + "litellm.llms.bedrock.base_aws_llm", + "BaseAWSLLM", + "_get_aws_region_name", + "get_aws_region_name", + "instance", + False, + ), + ( + "litellm.llms.bedrock.base_aws_llm", + "BaseAWSLLM", + "_validate_aws_region_name", + "validate_aws_region_name", + "staticmethod", + False, + ), + ( + "litellm.llms.bedrock.batches.transformation", + "BedrockBatchesConfig", + "_parse_timestamps_and_status", + "parse_timestamps_and_status", + "instance", + False, + ), + ( + "litellm.llms.bedrock.chat.agentcore.transformation", + "AmazonAgentCoreConfig", + "_get_runtime_session_id", + "get_runtime_session_id", + "instance", + False, + ), + ( + "litellm.llms.bedrock.chat.agentcore.transformation", + "AmazonAgentCoreConfig", + "_get_runtime_user_id", + "get_runtime_user_id", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_nova_transformation", + "AmazonNovaEmbeddingConfig", + "_transform_async_invoke_response", + "transform_async_invoke_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_nova_transformation", + "AmazonNovaEmbeddingConfig", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_nova_transformation", + "AmazonNovaEmbeddingConfig", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_titan_g1_transformation", + "AmazonTitanG1Config", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_titan_g1_transformation", + "AmazonTitanG1Config", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_titan_multimodal_transformation", + "AmazonTitanMultimodalEmbeddingG1Config", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_titan_multimodal_transformation", + "AmazonTitanMultimodalEmbeddingG1Config", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_titan_v2_transformation", + "AmazonTitanV2Config", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.amazon_titan_v2_transformation", + "AmazonTitanV2Config", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.cohere_transformation", + "BedrockCohereEmbeddingConfig", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.embedding", + "BedrockEmbedding", + "_get_async_invoke_status", + "get_async_invoke_status", + "instance", + True, + ), + ( + "litellm.llms.bedrock.embed.twelvelabs_marengo_transformation", + "TwelveLabsMarengoEmbeddingConfig", + "_transform_async_invoke_response", + "transform_async_invoke_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.twelvelabs_marengo_transformation", + "TwelveLabsMarengoEmbeddingConfig", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.bedrock.embed.twelvelabs_marengo_transformation", + "TwelveLabsMarengoEmbeddingConfig", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock.files.transformation", + "BedrockFilesConfig", + "_transform_openai_jsonl_content_to_bedrock_jsonl_content", + "transform_openai_jsonl_content_to_bedrock_jsonl_content", + "instance", + False, + ), + ( + "litellm.llms.bedrock.image_edit.amazon_nova_canvas_image_edit_transformation", + "BedrockAmazonNovaCanvasImageEditConfig", + "_is_nova_canvas_image_edit_model", + "is_nova_canvas_image_edit_model", + "classmethod", + False, + ), + ( + "litellm.llms.bedrock.image_edit.stability_transformation", + "BedrockStabilityImageEditConfig", + "_is_stability_edit_model", + "is_stability_edit_model", + "classmethod", + False, + ), + ( + "litellm.llms.bedrock.image_generation.amazon_nova_canvas_transformation", + "AmazonNovaCanvasConfig", + "_is_nova_model", + "is_nova_model", + "classmethod", + False, + ), + ( + "litellm.llms.bedrock.image_generation.amazon_stability3_transformation", + "AmazonStability3Config", + "_is_stability_3_model", + "is_stability_3_model", + "classmethod", + False, + ), + ( + "litellm.llms.bedrock.image_generation.amazon_titan_transformation", + "AmazonTitanImageGenerationConfig", + "_is_titan_model", + "is_titan_model", + "classmethod", + False, + ), + ( + "litellm.llms.bedrock.rerank.transformation", + "BedrockRerankConfig", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.bedrock.rerank.transformation", + "BedrockRerankConfig", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.bedrock_mantle.chat.transformation", + "BedrockMantleChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.chatgpt.chat.transformation", + "ChatGPTConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.clarifai.chat.transformation", + "ClarifaiConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.codestral.completion.transformation", + "CodestralTextCompletionConfig", + "_chunk_parser", + "chunk_parser", + "instance", + False, + ), + ( + "litellm.llms.cohere.embed.v1_transformation", + "CohereEmbeddingConfig", + "_populate_embedding_response", + "populate_embedding_response", + "instance", + False, + ), + ( + "litellm.llms.cohere.embed.v1_transformation", + "CohereEmbeddingConfig", + "_transform_request", + "transform_request", + "instance", + False, + ), + ( + "litellm.llms.cohere.embed.v1_transformation", + "CohereEmbeddingConfig", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.custom_httpx.llm_http_handler", + "BaseLLMHTTPHandler", + "_handle_error", + "handle_error", + "instance", + False, + ), + ( + "litellm.llms.dashscope.chat.transformation", + "DashScopeChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.datarobot.chat.transformation", + "DataRobotConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.deepinfra.chat.transformation", + "DeepInfraConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.docker_model_runner.chat.transformation", + "DockerModelRunnerChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.featherless_ai.chat.transformation", + "FeatherlessAIConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.fireworks_ai.chat.transformation", + "FireworksAIConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.gemini.chat.transformation", + "GoogleAIStudioGeminiConfig", + "_transform_messages", + "transform_messages", + "instance", + False, + ), + ( + "litellm.llms.github_copilot.chat.transformation", + "GithubCopilotConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.gradient_ai.chat.transformation", + "GradientAIConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.groq.chat.transformation", + "GroqChatConfig", + "_transform_messages", + "transform_messages", + "instance", + False, + ), + ( + "litellm.llms.groq.chat.transformation", + "GroqChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.heroku.chat.transformation", + "HerokuChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.hosted_vllm.chat.transformation", + "HostedVLLMChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.hyperbolic.chat.transformation", + "HyperbolicChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.inception.chat.transformation", + "InceptionChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.jina_ai.embedding.transformation", + "JinaAIEmbeddingConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.lambda_ai.chat.transformation", + "LambdaAIChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.langflow.chat.transformation", + "LangFlowConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.langgraph.chat.transformation", + "LangGraphConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.lemonade.chat.transformation", + "LemonadeChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.litellm_proxy.chat.transformation", + "LiteLLMProxyChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.llamafile.chat.transformation", + "LlamafileChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.lm_studio.chat.transformation", + "LMStudioChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.mistral.chat.transformation", + "MistralConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.mistral.chat.transformation", + "MistralConfig", + "_transform_messages", + "transform_messages", + "instance", + False, + ), + ( + "litellm.llms.modelscope.chat.transformation", + "ModelScopeChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.moonshot.chat.transformation", + "MoonshotChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.morph.chat.transformation", + "MorphChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.nscale.chat.transformation", + "NscaleConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.openai.chat.gpt_5_transformation", + "OpenAIGPT5Config", + "_supports_reasoning_effort_level", + "supports_reasoning_effort_level", + "classmethod", + False, + ), + ( + "litellm.llms.openai.chat.gpt_transformation", + "OpenAIGPTConfig", + "_transform_messages", + "transform_messages", + "instance", + False, + ), + ( + "litellm.llms.openai.common_utils", + "BaseOpenAILLM", + "_get_async_http_client", + "get_async_http_client", + "staticmethod", + False, + ), + ( + "litellm.llms.openai.openai", + "OpenAIChatCompletion", + "_get_openai_client", + "get_openai_client", + "instance", + False, + ), + ( + "litellm.llms.openai.openai", + "OpenAIConfig", + "_transform_messages", + "transform_messages", + "instance", + False, + ), + ( + "litellm.llms.openai.realtime.handler", + "OpenAIRealtime", + "_construct_url", + "construct_url", + "instance", + False, + ), + ( + "litellm.llms.openai_like.chat.transformation", + "OpenAILikeChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.openai_like.dynamic_config", + "JSONProviderConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.perplexity.chat.transformation", + "PerplexityChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.ragflow.chat.transformation", + "RAGFlowConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.snowflake.utils", + "SnowflakeBaseConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.soniox.audio_transcription.transformation", + "SonioxAudioTranscriptionConfig", + "_build_response_from_payload", + "build_response_from_payload", + "instance", + False, + ), + ( + "litellm.llms.together_ai.rerank.transformation", + "TogetherAIRerankConfig", + "_transform_response", + "transform_response", + "instance", + False, + ), + ( + "litellm.llms.v0.chat.transformation", + "V0ChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.vercel_ai_gateway.chat.transformation", + "VercelAIGatewayConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_calculate_usage", + "calculate_usage", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_calculate_web_search_requests", + "calculate_web_search_requests", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_check_finish_reason", + "check_finish_reason", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_check_prompt_level_content_filter", + "check_prompt_level_content_filter", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_drop_search_tools_mixed_with_functions", + "drop_search_tools_mixed_with_functions", + "classmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_forward_gemini_function_call_id", + "forward_gemini_function_call_id", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_is_gemini_3_or_newer", + "is_gemini_3_or_newer", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_is_model_gemini_spec_model", + "is_model_gemini_spec_model", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_map_audio_params", + "map_audio_params", + "instance", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_map_function", + "map_function", + "instance", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_map_web_search_options", + "map_web_search_options", + "instance", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_process_candidates", + "process_candidates", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_search_tool_keys", + "search_tool_keys", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_set_grounding_usage_counters", + "set_grounding_usage_counters", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_set_stream_metadata_on_response", + "set_stream_metadata_on_response", + "staticmethod", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_transform_google_generate_content_to_openai_model_response", + "transform_google_generate_content_to_openai_model_response", + "instance", + False, + ), + ( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini", + "VertexGeminiConfig", + "_transform_messages", + "transform_messages", + "instance", + False, + ), + ( + "litellm.llms.vertex_ai.vertex_llm_base", + "VertexBase", + "_ensure_access_token", + "ensure_access_token", + "instance", + False, + ), + ( + "litellm.llms.vertex_ai.vertex_llm_base", + "VertexBase", + "_ensure_access_token_async", + "ensure_access_token_async", + "instance", + True, + ), + ( + "litellm.llms.vertex_ai.vertex_llm_base", + "VertexBase", + "_get_token_and_url", + "get_token_and_url", + "instance", + False, + ), + ( + "litellm.llms.watsonx.common_utils", + "IBMWatsonXMixin", + "_prepare_payload", + "prepare_payload", + "instance", + False, + ), + ( + "litellm.llms.xai.chat.transformation", + "XAIChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), + ( + "litellm.llms.xai.chat.transformation", + "XAIChatConfig", + "_normalize_openai_compatible_usage_totals", + "normalize_openai_compatible_usage_totals", + "staticmethod", + False, + ), + ( + "litellm.llms.xai.chat.transformation", + "XAIChatConfig", + "_supports_stop_reason", + "supports_stop_reason", + "instance", + False, + ), + ( + "litellm.llms.zai.chat.transformation", + "ZAIChatConfig", + "_get_openai_compatible_provider_info", + "get_openai_compatible_provider_info", + "instance", + False, + ), +) def _get_owner(module_path: str, owner_name: str) -> object: module: Final = import_module(module_path or "litellm") + if module_path == "litellm.llms.openai_like.dynamic_config" and owner_name == "JSONProviderConfig": + provider_registry: Final = getattr( + import_module("litellm.llms.openai_like.json_loader"), + "JSONProviderRegistry", + ) + provider: Final = provider_registry.get("publicai") + return getattr(module, "create_config_class")(provider) return module if not owner_name else getattr(module, owner_name) @@ -989,7 +2242,7 @@ def _get_instance(owner: object, public_name: str) -> object: @pytest.mark.parametrize( ("module_path", "owner_name", "old_name", "new_name", "use_instance"), - ALIAS_CASES, + (*ALIAS_CASES, *LLMS_ALIAS_CASES), ) def test_public_aliases( module_path: str, @@ -1007,6 +2260,105 @@ def test_public_aliases( assert old_value is new_value +def _make_private_override( + descriptor: str, + is_async: bool, + result: object, +) -> object: + def recorded_args(args: tuple[object, ...]) -> tuple[object, ...]: + return args[1:] if descriptor in ("instance", "classmethod") else args + + def sync_override( + *args: object, **kwargs: object + ) -> tuple[object, tuple[object, ...], dict[str, object]]: + return result, recorded_args(args), kwargs + + async def async_override( + *args: object, **kwargs: object + ) -> tuple[object, tuple[object, ...], dict[str, object]]: + return result, recorded_args(args), kwargs + + implementation: Final = async_override if is_async else sync_override + if descriptor == "classmethod": + return classmethod(implementation) + if descriptor == "staticmethod": + return staticmethod(implementation) + return implementation + + +def _abstract_method_placeholder(*args: object, **kwargs: object) -> None: + return None + + +def _forwarder_arguments( + method: Callable[..., object], +) -> tuple[tuple[object, ...], dict[str, object]]: + parameters: Final = tuple(inspect.signature(method).parameters.values()) + positional: Final = tuple( + object() + for parameter in parameters + if parameter.kind + in ( + inspect.Parameter.POSITIONAL_ONLY, + inspect.Parameter.POSITIONAL_OR_KEYWORD, + inspect.Parameter.VAR_POSITIONAL, + ) + ) + keyword_only: Final = { + parameter.name: object() for parameter in parameters if parameter.kind is inspect.Parameter.KEYWORD_ONLY + } + variadic_keyword: Final = { + f"forwarder_extra_{index}": object() + for index, parameter in enumerate(parameters) + if parameter.kind is inspect.Parameter.VAR_KEYWORD + } + return positional, {**keyword_only, **variadic_keyword} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ( + "module_path", + "owner_name", + "private_name", + "public_name", + "descriptor", + "is_async", + ), + LLMS_FORWARDER_CASES, +) +async def test_public_forwarders_dispatch_to_private_subclass_override( + module_path: str, + owner_name: str, + private_name: str, + public_name: str, + descriptor: str, + is_async: bool, +) -> None: + owner: Final = _get_owner(module_path, owner_name) + if not isinstance(owner, type): + raise TypeError(f"expected a class owner, got {type(owner).__name__}") + expected: Final = object() + abstract_method_overrides: Final = { + name: _abstract_method_placeholder for name in getattr(owner, "__abstractmethods__", ()) + } + subclass: Final = type( + f"{owner_name}PrivateOverride", + (owner,), + {**abstract_method_overrides, private_name: _make_private_override(descriptor, is_async, expected)}, + ) + instance: Final = object.__new__(subclass) + public_method: Final = cast(Callable[..., object], getattr(instance, public_name)) + args, kwargs = _forwarder_arguments(public_method) + invocation: Final = public_method(*args, **kwargs) + if is_async: + assert inspect.isawaitable(invocation) + awaited: Final = await cast(Awaitable[object], invocation) + assert awaited == (expected, args, kwargs) + return + assert invocation == (expected, args, kwargs) + + @pytest.mark.parametrize( ("module_path", "owner_name", "old_name", "new_name"), PROPERTY_CASES, diff --git a/tests/unit/test_retrieve_batch_bedrock_dispatch.py b/tests/unit/test_retrieve_batch_bedrock_dispatch.py index 9d0daa52645..059ff430a4a 100644 --- a/tests/unit/test_retrieve_batch_bedrock_dispatch.py +++ b/tests/unit/test_retrieve_batch_bedrock_dispatch.py @@ -3,8 +3,8 @@ The dispatch picks one of two Bedrock handlers depending on the ARN family in ``batch_id``: -* ``:async-invoke/`` -> ``_handle_async_invoke_status`` (data plane) -* ``:model-invocation-job/`` -> ``_handle_model_invocation_job_status`` +* ``:async-invoke/`` -> ``handle_async_invoke_status`` (data plane) +* ``:model-invocation-job/`` -> ``handle_model_invocation_job_status`` (control plane, added in this PR) Anything else falls through to the generic ``provider_config`` retrieve @@ -37,11 +37,11 @@ def mock_handlers(): fake_batch = MagicMock(name="LiteLLMBatch") with ( patch( - "litellm.batches.main.BedrockBatchesHandler._handle_async_invoke_status", + "litellm.batches.main.BedrockBatchesHandler.handle_async_invoke_status", return_value=fake_batch, ) as async_invoke, patch( - "litellm.batches.main.BedrockBatchesHandler._handle_model_invocation_job_status", + "litellm.batches.main.BedrockBatchesHandler.handle_model_invocation_job_status", return_value=fake_batch, ) as mij, ): diff --git a/tests/unit/test_video_generation.py b/tests/unit/test_video_generation.py index 40fb6be67fe..dc15e82597d 100644 --- a/tests/unit/test_video_generation.py +++ b/tests/unit/test_video_generation.py @@ -553,7 +553,7 @@ class TestVideoGeneration: mock_client.post.return_value = mock_response with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ): result = handler.video_generation_handler( @@ -1113,7 +1113,7 @@ def test_video_content_handler_passes_variant_to_url(): mock_client.get.return_value = mock_response with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ): result = handler.video_content_handler( @@ -1158,11 +1158,9 @@ def test_video_content_handler_uses_get_for_openai(): mock_response.status_code = 200 mock_client.get.return_value = mock_response - # Patch _get_httpx_client to ensure no real HTTP client is created + # Patch get_httpx_client to ensure no real HTTP client is created # This prevents test isolation issues where isinstance check might fail - with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client" - ) as mock_get_client: + with patch("litellm.llms.custom_httpx.llm_http_handler.get_httpx_client") as mock_get_client: mock_get_client.return_value = mock_client result = handler.video_content_handler( @@ -1789,7 +1787,7 @@ def test_video_remix_handler_uses_api_key_from_litellm_params(): mock_client.post.return_value = MagicMock(status_code=200) with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ): handler.video_remix_handler( @@ -1876,7 +1874,7 @@ def test_video_remix_handler_prefers_explicit_api_key(): mock_client.post.return_value = MagicMock(status_code=200) with patch( - "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + "litellm.llms.custom_httpx.llm_http_handler.get_httpx_client", return_value=mock_client, ): handler.video_remix_handler(