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(