chore: merge litellm_internal_staging and resolve test conflicts

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
milan 2026-08-21 07:06:13 +00:00
commit 1cef8823fa
165 changed files with 2002 additions and 1018 deletions

View file

@ -191,7 +191,7 @@ class CustomStreamWrapper:
custom_llm_provider: str | None = None,
stream_options=None,
make_call: Callable | None = None,
_response_headers: dict | None = None,
_response_headers: dict | httpx.Headers | None = None,
):
self.model = model
self.make_call = make_call

View file

@ -35,7 +35,7 @@ def make_sync_call(
json_mode: bool | None = False,
fake_stream: bool = False,
stream_chunk_size: int | None = None,
):
) -> tuple[Any, httpx.Headers]:
if client is None:
client = _get_httpx_client() # Create a new client if none provided
@ -76,7 +76,7 @@ def make_sync_call(
additional_args={"complete_input_dict": data},
)
return completion_stream
return completion_stream, response.headers
class BedrockConverseLLM(BaseAWSLLM):
@ -134,7 +134,7 @@ class BedrockConverseLLM(BaseAWSLLM):
},
)
completion_stream: Final = await make_call(
completion_stream, response_headers = await make_call(
client=client,
api_base=api_base,
headers=dict(prepped.headers),
@ -151,6 +151,7 @@ class BedrockConverseLLM(BaseAWSLLM):
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
_response_headers=response_headers,
)
return streaming_response
@ -232,7 +233,7 @@ class BedrockConverseLLM(BaseAWSLLM):
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
return litellm.AmazonConverseConfig()._transform_response(
transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
model=model,
response=response,
model_response=model_response,
@ -244,6 +245,8 @@ class BedrockConverseLLM(BaseAWSLLM):
optional_params=optional_params,
encoding=encoding,
)
transformed_response.set_provider_response_headers(response.headers)
return transformed_response
def completion(
self,
@ -541,7 +544,7 @@ class BedrockConverseLLM(BaseAWSLLM):
client = client
if stream is not None and stream is True:
completion_stream: Final = make_sync_call(
completion_stream, response_headers = make_sync_call(
client=(client if client is not None and isinstance(client, HTTPHandler) else None),
api_base=proxy_endpoint_url,
headers=prepped.headers,
@ -558,6 +561,7 @@ class BedrockConverseLLM(BaseAWSLLM):
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
_response_headers=response_headers,
)
return streaming_response
@ -578,7 +582,7 @@ class BedrockConverseLLM(BaseAWSLLM):
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
return litellm.AmazonConverseConfig()._transform_response(
sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
model=model,
response=response,
model_response=model_response,
@ -590,3 +594,5 @@ class BedrockConverseLLM(BaseAWSLLM):
optional_params=optional_params,
encoding=encoding,
)
sync_transformed_response.set_provider_response_headers(response.headers)
return sync_transformed_response

View file

@ -163,7 +163,7 @@ async def make_call(
json_mode: bool | None = False,
bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None = None,
stream_chunk_size: int | None = None,
):
) -> tuple[Any, httpx.Headers]:
try:
if client is None:
client = get_async_httpx_client(
@ -225,7 +225,7 @@ async def make_call(
additional_args={"complete_input_dict": data},
)
return completion_stream
return completion_stream, response.headers
except httpx.HTTPStatusError as err:
error_code: Final = err.response.status_code
raise BedrockError(status_code=error_code, message=err.response.text)
@ -248,7 +248,7 @@ def make_sync_call(
json_mode: bool | None = False,
bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None = None,
stream_chunk_size: int | None = None,
):
) -> tuple[Any, httpx.Headers]:
try:
if client is None:
client = _get_httpx_client(
@ -309,7 +309,7 @@ def make_sync_call(
additional_args={"complete_input_dict": data},
)
return completion_stream
return completion_stream, response.headers
except httpx.HTTPStatusError as err:
error_code: Final = err.response.status_code
raise BedrockError(status_code=error_code, message=err.response.text)

View file

@ -1,7 +1,6 @@
import copy
import json
import time
from functools import partial
from typing import TYPE_CHECKING, Any, Final, cast, get_args
import httpx
@ -446,24 +445,24 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
) -> CustomStreamWrapper:
completion_stream, response_headers = await make_call(
client=client,
api_base=api_base,
headers=headers,
data=json.dumps(data),
model=model,
messages=messages,
logging_obj=logging_obj,
fake_stream=True if "ai21" in api_base else False,
bedrock_invoke_provider=self.get_bedrock_invoke_provider(model),
json_mode=json_mode,
)
streaming_response: Final = CustomStreamWrapper(
completion_stream=None,
make_call=partial(
make_call,
client=client,
api_base=api_base,
headers=headers,
data=json.dumps(data),
model=model,
messages=messages,
logging_obj=logging_obj,
fake_stream=True if "ai21" in api_base else False,
bedrock_invoke_provider=self.get_bedrock_invoke_provider(model),
json_mode=json_mode,
),
completion_stream=completion_stream,
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
_response_headers=response_headers,
)
return streaming_response
@ -481,27 +480,28 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
json_mode: bool | None = None,
signed_json_body: bytes | None = None,
) -> CustomStreamWrapper:
if client is None or isinstance(client, AsyncHTTPHandler):
client = _get_httpx_client(params={})
sync_client: Final = (
_get_httpx_client(params={}) if client is None or isinstance(client, AsyncHTTPHandler) else client
)
completion_stream, response_headers = make_sync_call(
client=sync_client,
api_base=api_base,
headers=headers,
data=json.dumps(data),
signed_json_body=signed_json_body,
model=model,
messages=messages,
logging_obj=logging_obj,
fake_stream=True if "ai21" in api_base else False,
bedrock_invoke_provider=self.get_bedrock_invoke_provider(model),
json_mode=json_mode,
)
streaming_response: Final = CustomStreamWrapper(
completion_stream=None,
make_call=partial(
make_sync_call,
client=client,
api_base=api_base,
headers=headers,
data=json.dumps(data),
signed_json_body=signed_json_body,
model=model,
messages=messages,
logging_obj=logging_obj,
fake_stream=True if "ai21" in api_base else False,
bedrock_invoke_provider=self.get_bedrock_invoke_provider(model),
json_mode=json_mode,
),
completion_stream=completion_stream,
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
_response_headers=response_headers,
)
return streaming_response

View file

@ -635,6 +635,7 @@ class BaseLLMHTTPHandler:
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
_response_headers=headers,
)
if client is None or not isinstance(client, HTTPHandler):
@ -798,6 +799,7 @@ class BaseLLMHTTPHandler:
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
_response_headers=_response_headers,
)
return streamwrapper

View file

@ -54,7 +54,30 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM):
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
return headers
inference_component_name: Final = optional_params.get("model_id")
if not isinstance(inference_component_name, str):
return headers
return {**headers, "X-Amzn-SageMaker-Inference-Component": inference_component_name}
def transform_request(
self,
model: str,
messages: list[AllMessageValues], # mutable-ok: matches the base chat transform signature
optional_params: dict, # mutable-ok: matches the base chat transform signature
litellm_params: dict, # mutable-ok: matches the base chat transform signature
headers: dict, # mutable-ok: matches the base chat transform signature
) -> dict: # mutable-ok: the handler sends this body straight to httpx
request: Final = super().transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
headers=headers,
)
served_model_name: Final = litellm_params.get("hf_model_name")
if not isinstance(served_model_name, str):
return request
return {**request, "model": served_model_name}
def get_complete_url(
self,

View file

@ -49489,6 +49489,16 @@
"source": "https://docs.devin.ai/windsurf/plugins/cascade/models"
},
"cognition/swe-1.7": {
"input_cost_per_token": 5e-07,
"output_cost_per_token": 2.5e-06,
"cache_read_input_token_cost": 2e-07,
"litellm_provider": "cognition",
"mode": "chat",
"supports_function_calling": true,
"supports_prompt_caching": true,
"source": "https://docs.devin.ai/desktop/models"
},
"cognition/swe-1.7-lightning": {
"input_cost_per_token": 2.5e-06,
"output_cost_per_token": 1.25e-05,
"cache_read_input_token_cost": 1e-06,
@ -49496,7 +49506,7 @@
"mode": "chat",
"supports_function_calling": true,
"supports_prompt_caching": true,
"source": "https://docs.devin.ai/windsurf/plugins/cascade/models"
"source": "https://docs.devin.ai/desktop/models"
},
"pinstripes/ps/glm-4.5-air": {
"max_tokens": 128000,

View file

@ -2640,51 +2640,63 @@ async def _process_team_members(
return updated_users, updated_team_memberships
def _resolve_member_identity(member: Member, updated_users: Sequence[LiteLLM_UserTable]) -> Member:
"""Return ``member`` with whichever of ``user_id`` / ``user_email`` the caller left out filled in.
The roster entry is a snapshot, so whatever is missing here is missing for good.
Resolution runs both ways off the user rows the add just touched: added by email
-> stamp the user_id, added by user_id -> stamp the email. A value the caller
supplied is never overwritten.
"""
resolved_user_id: Final = member.user_id or next(
(
user.user_id
for user in updated_users
if member.user_email is not None and user.user_email == member.user_email
),
None,
)
resolved_user_email: Final = member.user_email or next(
(
user.user_email
for user in updated_users
if resolved_user_id is not None and user.user_id == resolved_user_id and user.user_email is not None
),
None,
)
return member.model_copy(
update={ # mutable-ok: pydantic update payload
"user_id": resolved_user_id,
"user_email": resolved_user_email,
}
)
def _member_already_in_team(member: Member, complete_team_data: LiteLLM_TeamTable) -> bool:
return any(
(member.user_id is not None and existing_member.user_id == member.user_id)
or (member.user_email is not None and existing_member.user_email == member.user_email)
for existing_member in complete_team_data.members_with_roles
)
async def _update_team_members_list(
data: TeamMemberAddRequest,
complete_team_data: LiteLLM_TeamTable,
updated_users: list[LiteLLM_UserTable],
) -> None:
"""Update the team's members_with_roles list."""
if isinstance(data.member, Member):
new_member: Final = data.member.model_copy()
requested_members: Final[Sequence[Member]] = (
(data.member,) if isinstance(data.member, Member) else tuple(data.member)
)
resolved_members: Final = tuple(_resolve_member_identity(m, updated_users) for m in requested_members)
# get user id
if new_member.user_id is None and new_member.user_email is not None:
for user in updated_users:
if user.user_email is not None and user.user_email == new_member.user_email:
new_member.user_id = user.user_id
# Check if member already exists in team before adding
member_already_exists = False
for existing_member in complete_team_data.members_with_roles:
if (new_member.user_id is not None and existing_member.user_id == new_member.user_id) or (
new_member.user_email is not None and existing_member.user_email == new_member.user_email
):
member_already_exists = True
break
if not member_already_exists:
complete_team_data.members_with_roles.append(new_member)
elif isinstance(data.member, list):
for nm in data.member:
if nm.user_id is None and nm.user_email is not None:
for user in updated_users:
if user.user_email is not None and user.user_email == nm.user_email:
nm.user_id = user.user_id
# Check if member already exists in team before adding
member_already_exists = False
for existing_member in complete_team_data.members_with_roles:
if (nm.user_id is not None and existing_member.user_id == nm.user_id) or (
nm.user_email is not None and existing_member.user_email == nm.user_email
):
member_already_exists = True
break
if not member_already_exists:
complete_team_data.members_with_roles.append(nm)
# extend() consumes the generator as it appends, so a member already added by this
# same call is seen by the next _member_already_in_team check - the batch dedupes
# against itself exactly as the append-one-at-a-time loop this replaced did.
complete_team_data.members_with_roles.extend( # rebind-ok: this helper's contract is to grow the caller's roster in place
m for m in resolved_members if not _member_already_in_team(m, complete_team_data)
)
async def _add_team_members_to_team(
@ -4086,6 +4098,39 @@ async def _add_team_member_budget_table(
return team_info_response_object
async def _hydrate_member_emails(
prisma_client: PrismaClient,
members: Sequence[Member],
) -> tuple[Member, ...]:
"""Fill in ``user_email`` for roster entries that were stored without one.
``members_with_roles`` is a denormalized snapshot written at add-time, so an entry
stored with ``user_email=None`` keeps that null even once the user row has an email.
Look the missing ones up in ``LiteLLM_UserTable`` (one indexed query) and fill them
in. A stored email is never overwritten - the snapshot stays the source of truth
wherever it has a value.
"""
missing_user_ids: Final = frozenset(m.user_id for m in members if not m.user_email and m.user_id is not None)
if not missing_user_ids:
return tuple(members)
user_rows: Final[Sequence[LiteLLM_UserTable]] = await _user_db(prisma_client).find_many(
where={ # mutable-ok: Prisma query filters are dict-shaped
"user_id": { # mutable-ok: Prisma query filters are dict-shaped
"in": sorted(missing_user_ids)
}
}
)
email_by_user_id: Final = MappingProxyType({u.user_id: u.user_email for u in user_rows if u.user_email})
return tuple(
m.model_copy(update={"user_email": email_by_user_id[m.user_id]}) # mutable-ok: pydantic update payload
if not m.user_email and m.user_id in email_by_user_id
else m
for m in members
)
async def _resolve_team_access_group_resources(
_team_info: TeamInfoResponseObjectTeamTable,
) -> TeamInfoResponseObjectTeamTable:
@ -4221,9 +4266,22 @@ async def team_info(
# Resolve resources inherited from access groups
resolved_team_info: Final = await _resolve_team_access_group_resources(_team_info)
# Fill in emails the add-time roster snapshot never captured
hydrated_members: Final = await _hydrate_member_emails(
prisma_client=prisma_client,
members=resolved_team_info.members_with_roles,
)
hydrated_team_info: Final = resolved_team_info.model_copy(
update={ # mutable-ok: pydantic update payload
# list(), not the tuple: model_copy skips validation, so the field has
# to be handed the list[Member] the response model declares.
"members_with_roles": list(hydrated_members) # mutable-ok: declared list[Member]
}
)
response_object: Final = TeamInfoResponseObject(
team_id=team_id,
team_info=resolved_team_info,
team_info=hydrated_team_info,
keys=keys,
team_memberships=returned_tm,
)

View file

@ -10,6 +10,7 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
LITELLM_PROXY_MASTER_KEY_ALIAS,
LITELLM_TRUNCATED_PAYLOAD_FIELD,
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE,
REDACTED_BY_LITELM_STRING,
@ -21,6 +22,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
reconstruct_model_name,
)
from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
@ -53,13 +55,6 @@ def _get_max_string_length_prompt_in_db() -> int:
return DEFAULT_MAX_STRING_LENGTH_PROMPT_IN_DB
def _hash_api_key_for_spend_log(api_key: str) -> str:
stripped: Final = api_key[7:] if api_key[:7].lower() == "bearer " else api_key
if stripped.startswith("sk-"):
return hash_token(stripped)
return stripped
def _is_master_key(api_key: str | None, _master_key: str | None) -> bool:
"""
Raw-only constant-time master-key comparison. The hashed form is never
@ -70,6 +65,28 @@ def _is_master_key(api_key: str | None, _master_key: str | None) -> bool:
return secrets.compare_digest(api_key, _master_key)
_HASHED_JWT_RE = re.compile(r"hashed-jwt-[a-fA-F0-9]{64}")
def _is_non_secret_key_value(value: str) -> bool:
return (
value == LITELLM_PROXY_MASTER_KEY_ALIAS
or is_valid_sha256_hash(value)
or _HASHED_JWT_RE.fullmatch(value) is not None
)
def _redact_logged_api_key(value: str | None, *, already_redacted: bool = False) -> str | None:
if not isinstance(value, str) or not value:
return None
stripped: Final = re.sub(r"(?i)^bearer ", "", value)
if not stripped:
return None
if already_redacted and _is_non_secret_key_value(stripped):
return stripped
return hash_token(stripped)
def _get_spend_logs_metadata(
metadata: dict | None,
applied_guardrails: list[str] | None = None,
@ -123,9 +140,12 @@ def _get_spend_logs_metadata(
# Filter the metadata dictionary to include only the specified keys
clean_metadata: Final = SpendLogsMetadata(**{key: metadata.get(key) for key in SpendLogsMetadata.__annotations__})
raw_user_api_key: Final = clean_metadata.get("user_api_key")
if raw_user_api_key is not None and isinstance(raw_user_api_key, str):
clean_metadata["user_api_key"] = _hash_api_key_for_spend_log(raw_user_api_key)
_raw_key: Final = clean_metadata.get("user_api_key")
_trusted_hash: Final = metadata.get("user_api_key_hash")
_already_redacted: Final = (
isinstance(_trusted_hash, str) and _is_non_secret_key_value(_trusted_hash) and _trusted_hash == _raw_key
)
clean_metadata["user_api_key"] = _redact_logged_api_key(_raw_key, already_redacted=_already_redacted)
clean_metadata["applied_guardrails"] = applied_guardrails
clean_metadata["batch_models"] = batch_models
clean_metadata["mcp_tool_call_metadata"] = mcp_tool_call_metadata
@ -281,16 +301,23 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
standard_logging_prompt_tokens = standard_logging_payload.get("prompt_tokens", 0)
standard_logging_completion_tokens = standard_logging_payload.get("completion_tokens", 0)
standard_logging_total_tokens = standard_logging_payload.get("total_tokens", 0)
if api_key is not None and isinstance(api_key, str):
api_key = _hash_api_key_for_spend_log(api_key)
_trusted_hash = metadata.get("user_api_key_hash")
_key_already_redacted = (
isinstance(_trusted_hash, str) and _is_non_secret_key_value(_trusted_hash) and _trusted_hash == api_key
)
api_key = _redact_logged_api_key(api_key, already_redacted=_key_already_redacted) or ""
if (
standard_logging_payload is not None
): # [TODO] migrate completely to sl payload. currently missing pass-through endpoint data
api_key = api_key or standard_logging_payload["metadata"].get("user_api_key_hash") or ""
api_key = (
api_key
or _redact_logged_api_key(
standard_logging_payload["metadata"].get("user_api_key_hash"), already_redacted=True
)
or ""
)
end_user_id = end_user_id or standard_logging_payload["metadata"].get("user_api_key_end_user_id")
# BUG FIX: Don't overwrite api_key when standard_logging_payload is None
# The api_key was already extracted from metadata (line 243) and hashed (lines 256-259)
request_tags = safe_dumps(metadata.get("tags", [])) if isinstance(metadata.get("tags", []), list) else "[]"
if (
standard_logging_payload is not None and standard_logging_payload.get("request_tags") is not None

View file

@ -11,6 +11,7 @@ from typing import (
get_args,
)
import httpx
from openai._models import BaseModel as OpenAIObject
from openai.types.audio.transcription_create_params import (
FileTypes as FileTypes,
@ -49,7 +50,7 @@ from litellm.types.llms.base import (
)
from litellm.types.mcp import MCPServerCostInfo
from ..litellm_core_utils.core_helpers import map_finish_reason
from ..litellm_core_utils.core_helpers import map_finish_reason, process_response_headers
from .agents import LiteLLMSendMessageResponse
from .guardrails import GuardrailEventHooks
from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
@ -1916,6 +1917,10 @@ class ModelResponseBase(OpenAIObject):
_response_headers: dict | None = None
def set_provider_response_headers(self, headers: httpx.Headers) -> None:
"""Surface a provider's raw response headers to the caller as `llm_provider-*` headers."""
self._hidden_params["additional_headers"] = process_response_headers(headers)
def model_dump(self, **kwargs):
"""Default to exclude_unset to avoid Pydantic serializer warnings for OpenAIObject-derived types."""
if "exclude_unset" not in kwargs and "exclude_none" not in kwargs:

View file

@ -49489,6 +49489,16 @@
"source": "https://docs.devin.ai/windsurf/plugins/cascade/models"
},
"cognition/swe-1.7": {
"input_cost_per_token": 5e-07,
"output_cost_per_token": 2.5e-06,
"cache_read_input_token_cost": 2e-07,
"litellm_provider": "cognition",
"mode": "chat",
"supports_function_calling": true,
"supports_prompt_caching": true,
"source": "https://docs.devin.ai/desktop/models"
},
"cognition/swe-1.7-lightning": {
"input_cost_per_token": 2.5e-06,
"output_cost_per_token": 1.25e-05,
"cache_read_input_token_cost": 1e-06,
@ -49496,7 +49506,7 @@
"mode": "chat",
"supports_function_calling": true,
"supports_prompt_caching": true,
"source": "https://docs.devin.ai/windsurf/plugins/cascade/models"
"source": "https://docs.devin.ai/desktop/models"
},
"pinstripes/ps/glm-4.5-air": {
"max_tokens": 128000,

View file

@ -20,6 +20,12 @@
# PT012 a `pytest.raises` block that runs on past the raising call. Everything after
# that call is dead, so an `assert` sitting there is never checked. Keep the
# block to the call itself and put the assertions below it
# PT011 `pytest.raises(Exception)` / `(ValueError)` / `(OSError)` with no `match=`. The
# block passes on any error that broad, so the TypeError a refactor introduced
# reads as the rejection under test. Pin the message the code actually raises
# PT014 the same `parametrize` case listed twice. The copy re-runs an assertion that
# already passed and adds no coverage, and it usually marks a case someone meant
# to vary and forgot to edit
#
# No target-version here on purpose: it resolves from requires-python (>=3.10), so
# 3.11-only builtins like BaseExceptionGroup are correctly flagged in a tree that
@ -27,4 +33,4 @@
line-length = 120
lint.select = ["F821", "B011", "B015", "B017", "B018", "PT012", "PT015", "PLR0133", "PLW0127"]
lint.select = ["F821", "B011", "B015", "B017", "B018", "PT011", "PT012", "PT014", "PT015", "PLR0133", "PLW0127"]

View file

@ -400,7 +400,7 @@ def test_invalid_metric_name_validation():
litellm.prometheus_metrics_config = test_config
# Creating PrometheusLogger should raise ValueError due to invalid metric
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Configuration validation failed') as exc_info:
PrometheusLogger()
# Verify error message contains information about invalid metric
@ -429,7 +429,7 @@ def test_invalid_labels_validation():
litellm.prometheus_metrics_config = test_config
# Creating PrometheusLogger should raise ValueError due to invalid labels
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Configuration validation failed') as exc_info:
PrometheusLogger()
# Verify error message contains information about invalid labels
@ -598,7 +598,7 @@ def test_invalid_exclude_metric_name_raises(reset_prometheus_exclude_settings):
litellm.prometheus_exclude_labels = None
litellm.prometheus_exclude_metrics = ["not_a_real_metric"]
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Prometheus exclude configuration validation failed') as exc_info:
PrometheusLogger()
assert "not_a_real_metric" in str(exc_info.value)
@ -612,7 +612,7 @@ def test_invalid_exclude_label_name_raises(reset_prometheus_exclude_settings):
litellm.prometheus_exclude_metrics = None
litellm.prometheus_exclude_labels = ["not_a_real_label"]
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Prometheus exclude configuration validation failed') as exc_info:
PrometheusLogger()
assert "not_a_real_label" in str(exc_info.value)

View file

@ -141,7 +141,7 @@ async def test_bedrock_apply_guardrail_api_failure():
mock_api_request.side_effect = Exception("API connection failed")
# Test the apply_guardrail method should raise an exception
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Bedrock guardrail failed: API connection failed') as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["This is a test message"]},
request_data={},

View file

@ -1653,7 +1653,7 @@ async def test_afile_retrieve_raises_error_when_no_router_and_file_object_none()
unified_file_id = "test-unified-file-id"
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='LiteLLM Managed File object with id=test-unified-file-id') as exc_info:
await proxy_managed_files.afile_retrieve(
file_id=unified_file_id,
litellm_parent_otel_span=None,
@ -1719,7 +1719,7 @@ async def test_afile_retrieve_raises_error_for_non_managed_file():
# Mock get_unified_file_id to return None (file not found)
proxy_managed_files.get_unified_file_id = AsyncMock(return_value=None)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='LiteLLM Managed File object with id=non-existent-file-id') as exc_info:
await proxy_managed_files.afile_retrieve(
file_id="non-existent-file-id",
litellm_parent_otel_span=None,
@ -2027,7 +2027,7 @@ async def test_list_batches_from_managed_objects_table_provider_filter_raises_ex
)
# Filtering by provider should raise Exception
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Filtering by 'provider' is not supported when using managed") as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,
@ -2053,7 +2053,7 @@ async def test_list_batches_from_managed_objects_table_target_model_name_filter_
)
# Filtering by provider should raise Exception
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Filtering by 'target_model_names' is not supported when") as exc_info:
await proxy_managed_files.list_user_batches(
user_api_key_dict=UserAPIKeyAuth(user_id="test-user"),
limit=10,

View file

@ -448,7 +448,7 @@ def test_check_team_project_limits_models_not_in_team():
models=["gpt-5.5", "claude-3"], # claude-3 not in team
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="not in team's allowed models\\. Team allowed models") as exc_info:
_check_team_project_limits(team_object=team, data=data)
assert "claude-3" in str(exc_info.value.detail)
@ -476,7 +476,7 @@ def test_check_team_project_limits_budget_exceeds_team():
max_budget=150.0, # exceeds team's 100.0
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Project max_budget') as exc_info:
_check_team_project_limits(team_object=team, data=data)
assert "exceeds team's max_budget" in str(exc_info.value.detail)
@ -551,7 +551,7 @@ def test_check_team_project_limits_tpm_exceeds_team():
tpm_limit=20000, # exceeds team's 10000
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Project tpm_limit') as exc_info:
_check_team_project_limits(team_object=team, data=data)
assert "exceeds team's tpm_limit" in str(exc_info.value.detail)
@ -577,7 +577,7 @@ def test_check_team_project_limits_negative_budget():
max_budget=-10.0,
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='max_budget cannot be negative\\. Received') as exc_info:
_check_team_project_limits(team_object=team, data=data)
assert "cannot be negative" in str(exc_info.value.detail)
@ -604,7 +604,7 @@ def test_check_team_project_limits_soft_budget_gte_max():
soft_budget=100.0, # equal to max, should fail
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='must be strictly lower than max_budget') as exc_info:
_check_team_project_limits(team_object=team, data=data)
assert "must be strictly lower" in str(exc_info.value.detail)

View file

@ -61,7 +61,7 @@ async def test_dynamoai_blocks_content_with_block_action():
guardrail.should_run_guardrail = MagicMock(return_value=True)
# Test that the guardrail raises ValueError for blocked content
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='violation\\(s\\) detected') as exc_info:
await guardrail.async_pre_call_hook(
data=request_data,
user_api_key_dict=UserAPIKeyAuth(),

View file

@ -211,7 +211,7 @@ class TestEUAIActArticle5ConditionalMatching:
# Apply guardrail
if expected == "BLOCK":
# Should raise an exception or return modified response indicating block
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Content blocked: eu_ai_act_article') as exc_info:
await content_filter_guardrail.apply_guardrail(
inputs={"texts": [sentence]},
request_data=request_data,

View file

@ -83,7 +83,7 @@ class TestEUAIActFrench3Scenarios:
print(f"{'='*70}\n")
# Should raise an exception (blocked)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="prohibited_practices_fr conditional match 'concevoir \\+") as exc_info:
await content_filter_guardrail.apply_guardrail(
inputs={"texts": [sentence]},
request_data=request_data,
@ -123,7 +123,7 @@ class TestEUAIActFrench3Scenarios:
print(f"{'='*70}\n")
# Should raise an exception (blocked)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="prohibited_practices_fr conditional match 'créer \\+") as exc_info:
await content_filter_guardrail.apply_guardrail(
inputs={"texts": [sentence]},
request_data=request_data,
@ -194,7 +194,7 @@ class TestEUAIActFrench3Scenarios:
print(f"{'='*70}\n")
# Should raise an exception (blocked by conditional matching)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="prohibited_practices_fr conditional match 'développer \\+") as exc_info:
await content_filter_guardrail.apply_guardrail(
inputs={"texts": [sentence]},
request_data=request_data,
@ -278,7 +278,7 @@ class TestFrenchEdgeCases:
request_data = {"messages": [{"role": "user", "content": sentence}]}
# Should still block (no exception bypass)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="prohibited_practices_fr conditional match 'créer \\+ crédit") as exc_info:
await content_filter_guardrail.apply_guardrail(
inputs={"texts": [sentence]},
request_data=request_data,

View file

@ -55,7 +55,7 @@ def _make_guardrail(yaml_filename: str, category_name: str) -> ContentFilterGuar
async def _expect_block(guardrail: ContentFilterGuardrail, sentence: str, reason: str):
request_data = {"messages": [{"role": "user", "content": sentence}]}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Content blocked: sg_mas_') as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": [sentence]},
request_data=request_data,

View file

@ -62,7 +62,7 @@ def _make_guardrail(yaml_filename: str, category_name: str) -> ContentFilterGuar
async def _expect_block(guardrail: ContentFilterGuardrail, sentence: str, reason: str):
"""Assert that the guardrail BLOCKS the sentence."""
request_data = {"messages": [{"role": "user", "content": sentence}]}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Content blocked: sg_pdpa_') as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": [sentence]},
request_data=request_data,

View file

@ -432,7 +432,7 @@ def test_hashicorp_get_url_rejects_path_traversal(monkeypatch, malicious_secret_
monkeypatch.setenv("HCP_VAULT_TOKEN", "test-token-for-get-url-only")
manager = HashicorpSecretManager()
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='Invalid secret_name'):
manager.get_url(malicious_secret_name)

View file

@ -2147,7 +2147,7 @@ def test_validate_user_messages_invalid_content_type():
messages = [{"content": [{"type": "invalid_type", "text": "Hello"}]}]
with pytest.raises(Exception) as e:
with pytest.raises(Exception, match='Please ensure all messages are valid OpenAI chat completion') as e:
validate_chat_completion_user_messages(messages)
assert "Invalid message" in str(e)

View file

@ -37,27 +37,27 @@ def test_validate_tool_choice_cursor_format():
def test_validate_tool_choice_invalid_dict():
"""Test that invalid dict formats raise exceptions."""
# Missing both type and function
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Invalid tool choice, tool_choice=\\{\\}\\. Please ensure') as exc_info:
validate_chat_completion_tool_choice({})
assert "Invalid tool choice" in str(exc_info.value)
# Invalid type value
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Invalid tool choice, tool_choice=\\{'type': 'invalid'\\}\\.") as exc_info:
validate_chat_completion_tool_choice({"type": "invalid"})
assert "Invalid tool choice" in str(exc_info.value)
# Has type but missing function when type is "function"
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Invalid tool choice, tool_choice=\\{'type': 'function'\\}\\.") as exc_info:
validate_chat_completion_tool_choice({"type": "function"})
assert "Invalid tool choice" in str(exc_info.value)
def test_validate_tool_choice_invalid_type():
"""Test that invalid types raise exceptions."""
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="<class 'int'>\\. Expecting str, or dict\\. Please ensure") as exc_info:
validate_chat_completion_tool_choice(123)
assert "Got=<class 'int'>" in str(exc_info.value)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Invalid tool choice, tool_choice=\\[\\]\\. Got=<class 'list'>\\.") as exc_info:
validate_chat_completion_tool_choice([])
assert "Got=<class 'list'>" in str(exc_info.value)

View file

@ -295,7 +295,7 @@ async def test_responses_streaming_failure_triggers_failure_handlers():
call_type=CallTypes.responses.value,
)
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="boom"):
iterator._process_chunk('{"delta": "chunk"}')
# allow failure callbacks to run

View file

@ -1890,7 +1890,7 @@ def test_bedrock_completion_test_4(modify_params):
]
assert transformed_messages == expected_messages
else:
with pytest.raises(Exception) as e:
with pytest.raises(Exception, match=r"litellm\.modify_params") as e:
litellm.completion(**data)
assert "litellm.modify_params" in str(e.value)

View file

@ -12,6 +12,7 @@ This test suite verifies:
"""
from base_llm_unit_tests import BaseLLMChatTest
import httpx
import pytest
import sys
import os
@ -208,14 +209,6 @@ class TestBedrockMoonshotInvoke(BaseLLMChatTest):
endpoint with the messages body. Iteration of the stream itself is
not exercised here — moonshot streaming delegates to the OpenAI
parser and is covered by the OpenAI test suite.
Note: bedrock invoke streaming cannot be intercepted by patching
the caller-supplied client, because ``CustomStreamWrapper.fetch_sync_stream``
at streaming_handler.py invokes the stored ``make_call`` partial with
``client=litellm.module_level_client``, which overrides any client the
caller passed. Patch ``make_sync_call`` at its import site in
``base_invoke_transformation`` so we observe the exact kwargs the
partial was built with at stream-wrapper construction time.
"""
from litellm.utils import CustomStreamWrapper
@ -225,7 +218,7 @@ class TestBedrockMoonshotInvoke(BaseLLMChatTest):
captured.update(kwargs)
# Return an empty iterator so the stream wrapper's iteration
# doesn't try to parse real bytes.
return iter([])
return iter([]), httpx.Headers()
with patch(
"litellm.llms.bedrock.chat.invoke_transformations."
@ -246,11 +239,6 @@ class TestBedrockMoonshotInvoke(BaseLLMChatTest):
aws_region_name="us-west-2",
)
assert isinstance(response, CustomStreamWrapper)
# Trigger fetch_sync_stream → make_call(...) → fake_make_sync_call.
try:
next(iter(response))
except StopIteration:
pass
assert captured, "make_sync_call was never invoked"
assert captured["api_base"].endswith("/invoke-with-response-stream")

View file

@ -982,7 +982,7 @@ def test_convert_to_model_response_object_with_real_error():
},
}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception) as exc_info: # noqa: PT011 # message rides on .message, str() is empty
convert_to_model_response_object(
model_response_object=ModelResponse(),
response_object=response_object,
@ -1243,7 +1243,7 @@ def test_convert_to_model_response_object_with_error_code_only():
},
}
with pytest.raises(Exception) as exc_info: # noqa: B017 # bare Exception raised, so status_code is the assertion
with pytest.raises(Exception) as exc_info: # noqa: B017, PT011 # bare Exception, empty message, so status_code is the assertion
convert_to_model_response_object(
model_response_object=ModelResponse(),
response_object=response_object,
@ -1423,7 +1423,7 @@ def test_error_message_includes_function_args():
"choices": [{"index": 0}],
}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='in convert_to_model_response_object') as exc_info:
convert_to_model_response_object(
model_response_object=ModelResponse(),
response_object=response_object,

View file

@ -1845,7 +1845,7 @@ def test_parse_tool_call_arguments_malformed_json():
parse_tool_call_arguments,
)
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="Failed to parse tool call arguments for tool 'load_skill") as exc_info:
parse_tool_call_arguments(
'{"skill_name": "pptx',
tool_name="load_skill",
@ -1877,7 +1877,7 @@ def test_convert_to_anthropic_tool_invoke_malformed_json():
}
]
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="Failed to parse tool call arguments for tool 'bad_tool") as exc_info:
convert_to_anthropic_tool_invoke(tool_calls)
error_msg = str(exc_info.value)
@ -2023,7 +2023,7 @@ def test_parse_tool_call_arguments_still_raises_for_unrepairable():
parse_tool_call_arguments,
)
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="Failed to parse tool call arguments for tool 'test_tool") as exc_info:
parse_tool_call_arguments(
'{"key": "unterminated',
tool_name="test_tool",

View file

@ -45,7 +45,7 @@ def test_split_embedding_by_shape_fails_with_shape_value_error():
"data": [1, 2, 3, 4, 5, 6],
}
]
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='Shape must be of length'):
TritonEmbeddingConfig.split_embedding_by_shape(
data[0]["data"], data[0]["shape"]
)

View file

@ -59,7 +59,7 @@ def test_transform_request_invalid_provider(bedrock_transformer):
"""Test request transformation with invalid provider"""
messages = [{"role": "user", "content": "Hello"}]
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Bedrock Invoke HTTPX: Unknown provider=None') as exc_info:
bedrock_transformer.transform_request(
model="invalid.model",
messages=messages,

View file

@ -264,10 +264,6 @@ def test_get_end_user_id_from_request_body_backwards_compatibility():
["gpt-3.5-turbo", "gpt-4o-mini-general-deployment"],
),
({"model": "gpt-3.5-turbo"}, "gpt-3.5-turbo"),
(
{"model": "gpt-3.5-turbo, gpt-4o-mini-general-deployment"},
["gpt-3.5-turbo", "gpt-4o-mini-general-deployment"],
),
],
)
def test_get_model_from_request(request_data, expected_model):

View file

@ -1433,7 +1433,7 @@ async def test_exception_bubbling_up(sync_mode, stream_mode, model):
sync_stream=sync_mode,
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='litellm\\.BadRequestError: OpenAIException - Invalid value') as exc_info:
await _call_with_bad_role()
assert exc_info.value.code == "invalid_value"

View file

@ -23,13 +23,13 @@ class TestFileConsts:
def test_get_file_extension_from_mime_type(self):
assert get_file_extension_from_mime_type("audio/aac") == "aac"
assert get_file_extension_from_mime_type("application/pdf") == "pdf"
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='Unknown extension for mime type: application'):
get_file_extension_from_mime_type("application/unknown")
def test_get_file_type_from_extension(self):
assert get_file_type_from_extension("aac") == FileType.AAC
assert get_file_type_from_extension("pdf") == FileType.PDF
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='Unknown file type for extension: unknown'):
get_file_type_from_extension("unknown")
def test_get_file_extension_for_file_type(self):

View file

@ -134,7 +134,6 @@ def test_get_model_info_bedrock_region():
"ft:gpt-3.5-turbo:my-org:custom_suffix:id",
"ft:gpt-4-0613:my-org:custom_suffix:id",
"ft:davinci-002:my-org:custom_suffix:id",
"ft:gpt-4-0613:my-org:custom_suffix:id",
"ft:babbage-002:my-org:custom_suffix:id",
"gpt-35-turbo",
"ada",

View file

@ -160,7 +160,7 @@ async def test_provider_budgets_e2e_test_expect_to_fail():
await asyncio.sleep(2.5)
for _ in range(3):
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Exceeded budget for provider") as exc_info:
await router.acompletion(
messages=[{"role": "user", "content": "Hello, how are you?"}],
model="anthropic/claude-sonnet-4-5-20250929",
@ -594,7 +594,7 @@ async def test_deployment_budgets_e2e_test_expect_to_fail():
await asyncio.sleep(2.5)
for _ in range(3):
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Exceeded budget for deployment") as exc_info:
await router.acompletion(
messages=[{"role": "user", "content": "Hello, how are you?"}],
model="openai/gpt-4o-mini",
@ -646,7 +646,7 @@ async def test_tag_budgets_e2e_test_expect_to_fail():
await asyncio.sleep(2.5)
for _ in range(3):
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match=f"Exceeded budget for tag='{TAG_NAME}'") as exc_info:
await router.acompletion(
messages=[{"role": "user", "content": "Hello, how are you?"}],
model="openai/gpt-4o-mini",

View file

@ -1430,7 +1430,7 @@ async def test_router_fallbacks_default_and_model_specific_fallbacks(sync_mode):
messages=[{"role": "user", "content": "Hey, how's it going?"}],
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='litellm\\.AuthenticationError: AuthenticationError') as exc_info:
await _call_bad_model()
assert isinstance(
exc_info.value, litellm.AuthenticationError

View file

@ -293,7 +293,7 @@ def test_cleanup_timestamps():
assert all(isinstance(x, float) for x in result)
# Test invalid input
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="start_time is required, got=invalid of type <class 'str'>"):
StandardLoggingPayloadSetup.cleanup_timestamps(
"invalid", end_float, completion_float
)

View file

@ -143,7 +143,7 @@ async def test_team_blocking_behavior_multi_instance():
assert team_info_4001["blocked"] is True, "Team should be blocked after update"
# 8. Make a chat completion request on port 4000 with a new prompt; expect it to be blocked.
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match="(?i)blocked") as excinfo:
await chat_completion_on_port(
session,
key=key,
@ -157,7 +157,7 @@ async def test_team_blocking_behavior_multi_instance():
), f"Expected error indicating team blocked, got: {error_msg}"
# 9. Make a chat completion request on port 4000 with a new prompt; expect it to be blocked.
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match="(?i)blocked") as excinfo:
await chat_completion_on_port(
session,
key=key,
@ -171,7 +171,7 @@ async def test_team_blocking_behavior_multi_instance():
), f"Expected error indicating team blocked, got: {error_msg}"
# 9. Repeat the chat completion request with another new prompt; expect it to be blocked.
with pytest.raises(Exception) as excinfo_second:
with pytest.raises(Exception, match="(?i)blocked") as excinfo_second:
await chat_completion_on_port(
session,
key=key,

View file

@ -101,7 +101,7 @@ class TestAzureDocumentIntelligencePagesParam:
cfg.map_ocr_params({"pages": [True, False]}, {}, "prebuilt-layout")
def test_map_ocr_params_unsupported_type_raises(self, cfg):
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='based, Mistral-style\\) or a string like'):
cfg.map_ocr_params({"pages": 5}, {}, "prebuilt-layout")
def test_get_complete_url_appends_pages_query(self, cfg):

View file

@ -3,6 +3,7 @@ import asyncio
import aiohttp
import json
from httpx import AsyncClient
from openai import PermissionDeniedError
from typing import Any, Optional, List, Literal
@ -134,7 +135,7 @@ async def test_model_access_update():
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
# Should fail with gpt-5-mini
with pytest.raises(Exception) as exc_info:
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="openai/gpt-5-mini"
)
@ -157,7 +158,7 @@ async def test_model_access_update():
)
# Non-OpenAI model should still fail
with pytest.raises(Exception) as exc_info:
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="anthropic/claude-2"
)
@ -254,7 +255,7 @@ async def test_team_model_access_update():
await mock_chat_completion(session=session, key=key, model="openai/gpt-5.5")
# Should fail with gpt-5-mini
with pytest.raises(Exception) as exc_info:
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="openai/gpt-5-mini"
)
@ -279,7 +280,7 @@ async def test_team_model_access_update():
)
# Non-OpenAI model should still fail
with pytest.raises(Exception) as exc_info:
with pytest.raises(PermissionDeniedError) as exc_info:
await mock_chat_completion(
session=session, key=key, model="anthropic/claude-2"
)

View file

@ -1340,6 +1340,6 @@ async def test_team_model_alias(prisma_client, requested_model, should_pass):
}, "Expected model aliases to be present"
else:
# Verify the key fails with non-aliased models
with pytest.raises(Exception) as exc_info:
with pytest.raises(ProxyException) as exc_info:
await user_api_key_auth(request=request, api_key=f"Bearer {generated_key}")
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied

View file

@ -9,7 +9,7 @@ from litellm._uuid import uuid
from datetime import datetime
from dotenv import load_dotenv
from fastapi import Request
from fastapi import HTTPException, Request
from fastapi.routing import APIRoute
load_dotenv()
@ -530,7 +530,7 @@ async def test_user_role_permissions(prisma_client, route, user_role, expected_r
print(f"Auth passed as expected for {route} with role {user_role}")
else:
# Should raise an error
with pytest.raises(Exception) as exc_info:
with pytest.raises((ProxyException, HTTPException)) as exc_info:
await user_api_key_auth(request=request, api_key=bearer_token)
print(f"Auth failed as expected for {route} with role {user_role}")
print(f"Error message: {str(exc_info.value)}")

View file

@ -173,7 +173,7 @@ async def test_can_key_call_model(model, expect_to_work):
if expect_to_work:
await can_key_call_model(**args)
else:
with pytest.raises(Exception) as e:
with pytest.raises(Exception, match='key not allowed to access model\\. This key can only access') as e:
await can_key_call_model(**args)
print(e)
@ -958,7 +958,7 @@ async def test_can_key_call_model_with_aliases(model, alias_map, expect_to_work)
llm_router=router,
)
else:
with pytest.raises(Exception) as e:
with pytest.raises(Exception, match='key not allowed to access model\\. This key can only access') as e:
await can_key_call_model(
model=model,
llm_model_list=llm_model_list,

View file

@ -1583,7 +1583,7 @@ async def test_auth_jwt_mismatched_key_fails(monkeypatch):
h = JWTHandler()
with patch.object(h, "get_public_key", new=AsyncMock(return_value=rsa_jwk)):
with pytest.raises(Exception) as exc:
with pytest.raises(Exception, match='Validation fails: Expecting a PEM-formatted key\\.') as exc:
await h.auth_jwt(token)
assert "Validation fails" in str(exc.value)
@ -1826,7 +1826,7 @@ async def test_multi_issuer_jwt_unknown_issuer_without_global_jwks_rejected(
kid="issuer-key",
)
with pytest.raises(Exception) as exc:
with pytest.raises(Exception, match='Missing JWT Public Key URL from environment\\.') as exc:
await jwt_handler.auth_jwt(token=token)
assert "Missing JWT Public Key URL" in str(exc.value)
@ -1857,7 +1857,7 @@ async def test_multi_issuer_jwt_rejects_wrong_audience(monkeypatch):
kid="issuer-key",
)
with pytest.raises(Exception) as exc:
with pytest.raises(Exception, match="Validation fails: Audience doesn't match") as exc:
await jwt_handler.auth_jwt(token=token)
assert "Validation fails" in str(exc.value)
@ -1900,7 +1900,7 @@ async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch)
kid=shared_kid,
)
with pytest.raises(Exception) as exc:
with pytest.raises(Exception, match='Validation fails: Signature verification failed') as exc:
await jwt_handler.auth_jwt(token=token)
assert "Validation fails" in str(exc.value)
@ -1953,7 +1953,7 @@ def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled(
issuer = "https://issuer.example.com"
jwks_url = f"{issuer}/keys"
with pytest.raises(Exception) as exc:
with pytest.raises(Exception, match='must configure audience or set') as exc:
LiteLLM_JWTAuth(
issuers=[
{

View file

@ -1026,7 +1026,7 @@ def test_enforced_params_check(
from litellm.proxy.litellm_pre_call_utils import _enforced_params_check
if expected_error:
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='in request body\\. This is a required param'):
_enforced_params_check(
request_body=request_body,
general_settings=general_settings,
@ -2626,7 +2626,7 @@ async def test_during_call_hook_parallel_execution_with_error():
try:
litellm.callbacks = [FailingGuardrail()]
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Guardrail violation detected!') as exc_info:
await proxy_logging.during_call_hook(
data={
"model": "gpt-4",

View file

@ -166,7 +166,7 @@ async def test_update_spend_logs_non_connection_error():
prisma_client.db.litellm_spendlogs.create_many = create_many_mock
# Execute and verify it raises immediately without retrying
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Unexpected database error') as exc_info:
await update_spend(prisma_client, None, proxy_logging_obj)
# Verify error message

View file

@ -90,7 +90,7 @@ def test_routing_strategy_init_invalid_strategy(model_list):
router = Router(model_list=model_list)
# Test common mistake: "simple" instead of "simple-shuffle"
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="usage-based-routing', 'provider-budget-routing'\\]\\. Check") as exc_info:
router.routing_strategy_init(
routing_strategy="simple", routing_strategy_args={}
)
@ -106,7 +106,7 @@ def test_routing_strategy_init_invalid_strategy(model_list):
assert "Router SDK" in error_msg
# Test completely invalid strategy
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="usage-based-routing', 'provider-budget-routing'\\]\\. Check") as exc_info:
router.routing_strategy_init(
routing_strategy="not-a-real-strategy", routing_strategy_args={}
)

View file

@ -471,7 +471,7 @@ def test_validate_mcp_server_name_direct():
validate_mcp_server_name("valid name")
# Test that invalid names with hyphens raise exceptions
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="Server name cannot contain '-'\\. Use an alternative") as exc_info:
validate_mcp_server_name("invalid-name")
assert "cannot contain" in str(exc_info.value)

View file

@ -104,7 +104,7 @@ async def test_send_email_missing_api_key():
try:
logger = SendGridEmailLogger()
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='SENDGRID_API_KEY is not set'):
await logger.send_email(
from_email="test@example.com",
to_email=["recipient@example.com"],

View file

@ -471,7 +471,7 @@ async def test_afile_content_error_reports_unified_id_not_provider_uri():
mock_router.get_deployment_credentials_with_provider = MagicMock(return_value=None)
mock_router.afile_content = AsyncMock(side_effect=Exception("deployment failed"))
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='LiteLLM Managed File object with') as exc_info:
await managed_files.afile_content(
file_id=unified_file_id,
litellm_parent_otel_span=None,

View file

@ -51,7 +51,7 @@ async def test_get_usage_data_rejects_invalid_limit(monkeypatch: pytest.MonkeyPa
"""limit must coerce to int or raise ValueError before hitting the DB."""
db, query_mock = _setup_db(monkeypatch, [])
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='limit must be an integer'):
await db.get_usage_data(limit="invalid")
assert query_mock.await_count == 0

View file

@ -108,7 +108,7 @@ class TestCloudZeroStreamer:
"""Test _parse_and_convert_timestamp method with invalid timestamp."""
streamer = CloudZeroStreamer("test-key", "test-connection")
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="Could not parse timestamp 'invalid-timestamp': Invalid"):
streamer._parse_and_convert_timestamp("invalid-timestamp")
def test_prepare_batch_payload(self):

View file

@ -68,7 +68,7 @@ async def test_should_accept_string_timestamps(monkeypatch: pytest.MonkeyPatch):
async def test_should_reject_invalid_limit(monkeypatch: pytest.MonkeyPatch):
db, query_mock = _setup_db(monkeypatch, [])
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='limit must be an integer'):
await db.get_usage_data(limit="invalid")
assert query_mock.await_count == 0

View file

@ -20,7 +20,7 @@ def _window(freq: str = "hourly", hour: int = 5) -> FocusTimeWindow:
def test_should_require_bucket_name():
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='bucket_name must be provided for S'):
FocusS3Destination(prefix="focus", config={})

View file

@ -95,9 +95,9 @@ def enc_project(p): # how client encodes project in urls
# Constructor / config tests
# -----------------------------
def test_init_requires_project_and_token():
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='project and access_token are required'):
GitLabClient({"project": "p"})
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='project and access_token are required'):
GitLabClient({"access_token": "t"})
@ -127,7 +127,7 @@ def test_set_ref_updates_effective_ref():
c = make_client(branch="main")
c.set_ref("feature/x")
assert c.ref == "feature/x"
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='ref must be a non-empty string'):
c.set_ref("")
@ -193,12 +193,12 @@ def test_get_file_content_permission_errors_are_mapped():
raw_url = f"https://gitlab.example.com/api/v4/projects/{enc_project('group/sub/repo')}/repository/files/secure%2Ffile.prompt/raw?ref=main"
# raise_for_status will be called, so return 403 response (not an exception from transport)
c.http_handler.routes[raw_url] = FakeResponse(status_code=403)
with pytest.raises(Exception) as ei:
with pytest.raises(Exception, match="Check your GitLab permissions for project 'group") as ei:
c.get_file_content("secure/file.prompt")
assert "Access denied" in str(ei.value)
c.http_handler.routes[raw_url] = FakeResponse(status_code=401)
with pytest.raises(Exception) as ei2:
with pytest.raises(Exception, match='Authentication failed\\. Check your GitLab token and') as ei2:
c.get_file_content("secure/file.prompt")
assert "Authentication failed" in str(ei2.value)

View file

@ -198,7 +198,7 @@ class TestLevoIntegration(unittest.TestCase):
"""Test health check returns unhealthy status when required vars are missing."""
# Try to create logger without required env vars
# This should fail during config, but we can test health check logic
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='LEVOAI_API_KEY environment variable is required for Levo'):
LevoLogger.get_levo_config()
@patch.dict(

View file

@ -554,7 +554,7 @@ def test_token_type_rejected_from_either_list(attributes, monkeypatch):
recorder rather than silently ignored, so the misconfig is caught at all."""
recorder = _recorder(monkeypatch, attributes)
kwargs, response_obj, start, end = _build_call()
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='otel\\.attributes: gen_ai\\.token\\.type is a structural') as exc_info:
recorder.record(kwargs, response_obj, start, end)
# The dedicated discriminator guard, not the generic unknown-name path: assert
# the specific reason so dropping that guard (and falling through to "unknown

View file

@ -1163,7 +1163,7 @@ def test_max_langfuse_clients_limit():
assert litellm.initialized_langfuse_clients == 2
# Third client should fail with exception
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Max langfuse clients reached') as exc_info:
logger3 = LangFuseLogger(
langfuse_public_key="test_key_3",
langfuse_secret="test_secret_3",

View file

@ -1169,7 +1169,7 @@ def test_bedrock_image_processor_content_type_fallback_failure():
# Test with URL without recognizable extension
image_url = "https://example.com/unknown-file"
with pytest.raises(ValueError) as excinfo:
with pytest.raises(ValueError, match='Unable to determine content type from URL: https') as excinfo:
BedrockImageProcessor._post_call_image_processing(mock_response, image_url)
assert "Unable to determine content type" in str(excinfo.value)

View file

@ -110,7 +110,7 @@ def test_top_level_kwargs_overrides_metadata_slots():
def test_env_reference_at_top_level_raises_with_guidance():
kwargs = {"langfuse_public_key": "os.environ/LANGFUSE_PUBLIC_KEY"}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="Callback param 'langfuse_public_key' \\(from request body\\)") as exc_info:
initialize_standard_callback_dynamic_params(kwargs)
message = str(exc_info.value)
@ -127,7 +127,7 @@ def test_env_reference_in_metadata_raises_with_guidance():
}
}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="Callback param 'langsmith_api_key' \\(from metadata\\) contains") as exc_info:
initialize_standard_callback_dynamic_params(kwargs)
message = str(exc_info.value)

View file

@ -27,7 +27,7 @@ def test_parse_json_verdict_tolerates_fences_and_prose(raw, expected):
def test_parse_json_verdict_rejects_non_object():
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='judge response is not a JSON object'):
parse_json_verdict('["not", "an", "object"]')
with pytest.raises((json.JSONDecodeError, ValueError)):
parse_json_verdict("no json here at all")

View file

@ -982,7 +982,7 @@ async def test_bedrock_validation_error_raises_directly(logging_obj: Logging):
make_call=_raise_400,
)
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match='litellm\\.BadRequestError: BedrockException') as excinfo:
await response.__anext__()
assert not isinstance(excinfo.value, MidStreamFallbackError)
assert getattr(excinfo.value, "status_code", None) == 400
@ -2722,7 +2722,7 @@ def test_dispatch_text_completion_codestral_requires_string(
is a programming error and must surface loudly."""
initialized_custom_stream_wrapper.custom_llm_provider = "text-completion-codestral"
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="chunk is not a string: \\{'not': 'a string'\\}"):
_run_dispatch(initialized_custom_stream_wrapper, {"not": "a string"})

View file

@ -763,24 +763,6 @@ class TestTokenizerSelection(unittest.TestCase):
],
}
],
[
{
"role": "user",
"content": [
{
"type": "text",
"text": "These are some sample images from a movie. Based on these images, what do you think the tone of the movie is?",
},
{
"type": "text",
"image_url": {
"url": "https://gratisography.com/wp-content/uploads/2024/11/gratisography-augmented-reality-800x525.jpg",
"detail": "high",
},
},
],
}
],
],
)
def test_bad_input_token_counter(model, messages):
@ -1174,7 +1156,7 @@ def test_count_content_list_rejects_unknown_type():
"""
from litellm.litellm_core_utils.token_counter import _count_content_list
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Error getting number of tokens from content list: Invalid') as exc_info:
_count_content_list(
count_function=len,
content_list=[{"type": "totally_unknown_block"}],

View file

@ -100,12 +100,12 @@ class TestEncodeUrlPathSegment:
@pytest.mark.parametrize("value", ["", ".", "..", None])
def test_rejects_empty_and_dot_segments(self, value):
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="resource_id (is required|cannot be a dot path segment)"):
encode_url_path_segment(value, field_name="resource_id")
@pytest.mark.parametrize("value", ["../model", "model/../other", "/model"])
def test_rejects_dot_segments_in_multi_segment_paths(self, value):
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="model (is required|cannot be a dot path segment)"):
encode_url_path_segments(value, field_name="model")

View file

@ -113,7 +113,7 @@ def test_flux_style_request_still_remaps_to_legacy_fields():
def test_openai_style_unsupported_param_raises_without_drop_params():
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='Supported parameters are'):
AimlImageGenerationConfig().map_openai_params(
non_default_params={"image_size": {"width": 1024, "height": 1024}},
optional_params={},

View file

@ -61,7 +61,7 @@ def test_anthropic_messages_handler_skips_the_gateway_on_recursion():
"litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp",
new=AsyncMock(return_value={"routed": True}),
) as routed:
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='anthropic_messages_handler is not implemented for sync calls'):
anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],
@ -80,7 +80,7 @@ def test_anthropic_messages_handler_leaves_native_tools_alone():
"litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp",
new=AsyncMock(return_value={"routed": True}),
) as routed:
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='anthropic_messages_handler is not implemented for sync calls'):
anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "hi"}],

View file

@ -16,7 +16,7 @@ class TestAzureAIRerankConfigGetCompleteUrl:
self.model = "azure_ai/cohere-rerank-v3-english"
def test_api_base_required(self):
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Azure AI API Base is required\\. api_base=None\\. Set in') as exc_info:
self.config.get_complete_url(api_base=None, model=self.model)
assert "api_base=None" in str(exc_info.value)
@ -31,7 +31,7 @@ class TestAzureAIRerankConfigGetCompleteUrl:
],
)
def test_api_base_requires_scheme(self, api_base):
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Azure AI API Base must be an absolute URL including scheme') as exc_info:
self.config.get_complete_url(api_base=api_base, model=self.model)
error_message = str(exc_info.value).lower()

View file

@ -2,17 +2,20 @@ import os
import sys
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.bedrock.chat.invoke_handler import (
AWSEventStreamDecoder,
make_call,
make_sync_call,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
def test_transform_thinking_blocks_with_redacted_content():
@ -293,3 +296,50 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
response.iter_bytes.assert_called_once_with(chunk_size=2048)
def test_invoke_streaming_forwards_bedrock_response_headers():
response = MagicMock()
response.status_code = 200
response.iter_bytes = MagicMock(return_value=iter([]))
response.headers = httpx.Headers({"x-amzn-requestid": "req-789"})
client = HTTPHandler()
client.post = MagicMock(return_value=response)
stream = litellm.completion(
model="bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
)
assert stream._hidden_params["additional_headers"]["llm_provider-x-amzn-requestid"] == "req-789"
@pytest.mark.asyncio
async def test_async_invoke_streaming_forwards_bedrock_response_headers():
async def _no_bytes(chunk_size=None):
return
yield b""
response = MagicMock()
response.status_code = 200
response.aiter_bytes = _no_bytes
response.headers = httpx.Headers({"x-amzn-requestid": "req-987"})
client = AsyncHTTPHandler()
client.post = AsyncMock(return_value=response)
stream = await litellm.acompletion(
model="bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
)
assert stream._hidden_params["additional_headers"]["llm_provider-x-amzn-requestid"] == "req-987"

View file

@ -1944,7 +1944,7 @@ def test_role_assumption_access_denied_raises_when_different_role():
with patch.object(
base_aws_llm, "_is_already_running_as_role", return_value=False
):
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='An error occurred \\(AccessDenied\\) when calling the') as exc_info:
base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,
@ -1969,7 +1969,7 @@ def test_role_assumption_non_access_denied_error_propagated():
)
with patch("boto3.client", return_value=mock_sts_client):
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='An error occurred \\(MalformedPolicyDocument\\) when calling') as exc_info:
base_aws_llm._auth_with_aws_role(
aws_access_key_id=None,
aws_secret_access_key=None,

View file

@ -109,7 +109,7 @@ class TestBedrockMantleResponsesURL:
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
monkeypatch.delenv("AWS_REGION", raising=False)
cfg = BedrockMantleResponsesAPIConfig()
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="api\\.aws\\.attacker\\.example/'\\. Region names must contain only"):
cfg.get_complete_url(
api_base=None,
litellm_params={
@ -1418,7 +1418,7 @@ class TestBedrockMantleResponsesSigV4:
signer.get_credentials = MagicMock(side_effect=NoCredentialsError())
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
with pytest.raises(ValueError) as exc:
with pytest.raises(ValueError, match='Bedrock Mantle auth failed: no Bearer token and no usable') as exc:
cfg.sign_request(
headers={},
optional_params={"aws_region_name": "us-east-2"},
@ -1448,7 +1448,7 @@ class TestBedrockMantleResponsesSigV4:
signer.get_credentials = MagicMock(side_effect=cred_error)
cfg = BedrockMantleResponsesAPIConfig(aws_signer=signer)
with pytest.raises(ValueError) as exc:
with pytest.raises(ValueError, match='Bedrock Mantle auth failed: no Bearer token and no usable') as exc:
cfg.sign_request(
headers={},
optional_params={"aws_region_name": "us-east-2"},

View file

@ -107,7 +107,7 @@ class TestBedrockMantleConfig:
monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False)
monkeypatch.delenv("AWS_REGION", raising=False)
cfg = BedrockMantleChatConfig()
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="api\\.aws\\.attacker\\.example/'\\. Region names must contain only"):
cfg._get_openai_compatible_provider_info(
None,
None,
@ -416,7 +416,7 @@ class TestBedrockMantleChatAuth:
signer.get_credentials = MagicMock(side_effect=NoCredentialsError())
cfg = BedrockMantleChatConfig(aws_signer=signer)
with pytest.raises(ValueError) as exc:
with pytest.raises(ValueError, match='Bedrock Mantle auth failed: no Bearer token and no usable') as exc:
cfg.sign_request(
headers={},
optional_params={"aws_region_name": "us-east-2"},

View file

@ -38,7 +38,7 @@ class TestBytezChatConfig:
config = BytezChatConfig()
headers = {}
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match='Missing api_key, make sure you pass in your api key') as excinfo:
config.validate_environment(
headers=headers,
model=TEST_MODEL,

View file

@ -1,14 +1,16 @@
import json
import os
import sys
from unittest.mock import MagicMock
from unittest.mock import AsyncMock, MagicMock
import httpx
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.custom_httpx.http_handler import HTTPHandler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
sys.path.insert(
0, os.path.abspath("../../../../..")
@ -202,6 +204,104 @@ def test_make_sync_call_honors_explicit_stream_chunk_size():
response.iter_bytes.assert_called_once_with(chunk_size=2048)
def _converse_response_body() -> dict:
return {
"output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 1, "outputTokens": 1, "totalTokens": 2},
}
def test_converse_completion_forwards_bedrock_response_headers():
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json = MagicMock(return_value=_converse_response_body())
mock_response.text = json.dumps(_converse_response_body())
mock_response.headers = httpx.Headers({"x-amzn-requestid": "req-123"})
client = HTTPHandler()
client.post = MagicMock(return_value=mock_response)
response = litellm.completion(
model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
client=client,
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
)
assert response._hidden_params["additional_headers"]["llm_provider-x-amzn-requestid"] == "req-123"
def test_converse_streaming_forwards_bedrock_response_headers():
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.iter_bytes = MagicMock(return_value=iter([]))
mock_response.headers = httpx.Headers({"x-amzn-requestid": "req-456"})
client = HTTPHandler()
client.post = MagicMock(return_value=mock_response)
response = litellm.completion(
model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
)
assert response._hidden_params["additional_headers"]["llm_provider-x-amzn-requestid"] == "req-456"
@pytest.mark.asyncio
async def test_async_converse_completion_forwards_bedrock_response_headers():
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json = MagicMock(return_value=_converse_response_body())
mock_response.text = json.dumps(_converse_response_body())
mock_response.headers = httpx.Headers({"x-amzn-requestid": "req-abc"})
client = AsyncHTTPHandler()
client.post = AsyncMock(return_value=mock_response)
response = await litellm.acompletion(
model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
client=client,
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
)
assert response._hidden_params["additional_headers"]["llm_provider-x-amzn-requestid"] == "req-abc"
@pytest.mark.asyncio
async def test_async_converse_streaming_forwards_bedrock_response_headers():
async def _no_bytes(chunk_size=None):
return
yield b""
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.aiter_bytes = _no_bytes
mock_response.headers = httpx.Headers({"x-amzn-requestid": "req-def"})
client = AsyncHTTPHandler()
client.post = AsyncMock(return_value=mock_response)
response = await litellm.acompletion(
model="bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
aws_access_key_id="fake",
aws_secret_access_key="fake",
aws_region_name="us-east-1",
)
assert response._hidden_params["additional_headers"]["llm_provider-x-amzn-requestid"] == "req-def"
def test_completion_plumbs_stream_chunk_size_through_converse():
iter_bytes_spy = _stream_completion_with_spied_iter_bytes(
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0"

View file

@ -2369,3 +2369,79 @@ async def test_async_anthropic_messages_handler_carries_deployment_vertex_locati
unconfigured_deployment = await logging_obj_after_handler(GenericLiteLLMParams())
assert "vertex_location" not in unconfigured_deployment.litellm_params
_GENERIC_STREAM_SSE = (
b'data: {"id":"chatcmpl-1","object":"chat.completion.chunk","created":1,'
b'"model":"test-model","choices":[{"index":0,"delta":{"content":"hi"},'
b'"finish_reason":null}]}\n\n'
b"data: [DONE]\n\n"
)
def _generic_stream_upstream_response() -> httpx.Response:
return httpx.Response(
200,
headers={
"x-request-id": "generic-req-123",
"x-ratelimit-remaining-requests": "42",
},
content=_GENERIC_STREAM_SSE,
request=httpx.Request("POST", "https://fake-vllm.test/v1/chat/completions"),
)
def test_generic_http_handler_sync_streaming_forwards_provider_response_headers():
"""
Regression test for the generic BaseLLMHTTPHandler streaming path used by
~30 providers (deepseek, groq, hosted_vllm, databricks, openrouter, ...).
The sync `completion()` streaming branch builds the CustomStreamWrapper from
`make_sync_call`, which returns the upstream response headers alongside the
stream. Those headers must reach the caller as `llm_provider-*` entries in
`_hidden_params["additional_headers"]`, which is what the proxy merges into
the client-facing response headers.
"""
mock_client = Mock(spec=HTTPHandler)
mock_client.post = Mock(return_value=_generic_stream_upstream_response())
response = litellm.completion(
model="hosted_vllm/test-model",
messages=[{"role": "user", "content": "Hello"}],
api_base="https://fake-vllm.test/v1",
api_key="sk-test",
stream=True,
client=mock_client,
)
additional_headers = response._hidden_params["additional_headers"]
assert additional_headers["llm_provider-x-request-id"] == "generic-req-123"
assert additional_headers["llm_provider-x-ratelimit-remaining-requests"] == "42"
assert "".join([chunk.choices[0].delta.content or "" for chunk in response]) == "hi"
@pytest.mark.asyncio
async def test_generic_http_handler_async_streaming_forwards_provider_response_headers():
"""
Companion to the sync test above for `acompletion_stream_function`, which
builds its CustomStreamWrapper from `make_async_call_stream_helper`.
"""
mock_client = AsyncMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=_generic_stream_upstream_response())
response = await litellm.acompletion(
model="hosted_vllm/test-model",
messages=[{"role": "user", "content": "Hello"}],
api_base="https://fake-vllm.test/v1",
api_key="sk-test",
stream=True,
client=mock_client,
)
additional_headers = response._hidden_params["additional_headers"]
assert additional_headers["llm_provider-x-request-id"] == "generic-req-123"
assert additional_headers["llm_provider-x-ratelimit-remaining-requests"] == "42"
collected = [chunk async for chunk in response]
assert "".join([chunk.choices[0].delta.content or "" for chunk in collected]) == "hi"

View file

@ -258,7 +258,7 @@ class TestDeepinfraRerankTransform:
status_code = 401
headers = {"content-type": "application/json"}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Authentication failed') as exc_info:
self.config.get_error_class(error_message, status_code, headers)
# The method should raise a BaseLLMException
@ -271,7 +271,7 @@ class TestDeepinfraRerankTransform:
status_code = 404
headers = {"content-type": "application/json"}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Model not found') as exc_info:
self.config.get_error_class(error_message, status_code, headers)
# Should extract the nested error message
@ -284,7 +284,7 @@ class TestDeepinfraRerankTransform:
status_code = 503
headers = {"content-type": "application/json"}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Service unavailable') as exc_info:
self.config.get_error_class(error_message, status_code, headers)
# Should extract the string detail
@ -296,7 +296,7 @@ class TestDeepinfraRerankTransform:
status_code = 500
headers = {"content-type": "application/json"}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Invalid JSON error message') as exc_info:
self.config.get_error_class(error_message, status_code, headers)
# Should use the original error message when JSON parsing fails

View file

@ -113,7 +113,7 @@ def test_response_format_is_ignored():
def test_unsupported_param_raises_without_drop_params():
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="Supported parameters are \\['n', 'response_format', 'size'\\]\\."):
FalAINanoBananaConfig().map_openai_params(
non_default_params={"style": "vivid"},
optional_params={},

View file

@ -44,7 +44,7 @@ class TestFeatherlessAIConfig:
"""Test error handling when API key is missing"""
config = FeatherlessAIConfig()
with pytest.raises(ValueError) as excinfo:
with pytest.raises(ValueError, match='Missing Featherless AI API Key') as excinfo:
config.validate_environment(
headers={},
model="featherless-ai/Qwerky-72B",
@ -112,7 +112,7 @@ class TestFeatherlessAIConfig:
"tool_choice": {"type": "function", "function": {"name": "get_weather"}}
}
optional_params = {}
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match="litellm\\.UnsupportedParamsError: Featherless AI doesn't") as excinfo:
config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
@ -138,7 +138,7 @@ class TestFeatherlessAIConfig:
assert "tools" not in result
# Test with tools and drop_params=False
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match="litellm\\.UnsupportedParamsError: Featherless AI doesn't") as excinfo:
config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,

View file

@ -301,7 +301,7 @@ class TestFireworksAIRerankTransform:
mock_logging = MagicMock()
model_response = RerankResponse()
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Failed to parse response: Invalid JSON: line') as exc_info:
self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,

View file

@ -244,7 +244,7 @@ class TestGeminiImageEditTransformation:
def test_transform_image_edit_request_without_image_raises(self) -> None:
optional_params = {}
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='Gemini image edit requires at least one image\\.'):
self.config.transform_image_edit_request(
model=self.model,
prompt=self.prompt,

View file

@ -28,7 +28,7 @@ def test_gemini_completion_no_api_key():
del os.environ[key]
# Test without mock_response to ensure actual API key validation
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='in _complete_vertex_ai_beta') as exc_info:
completion(
model="gemini/gemini-1.5-flash",
messages=[{"role": "user", "content": "Test message"}],
@ -60,7 +60,7 @@ def test_gemini_completion_no_api_key_with_mock():
with patch("litellm.get_secret") as mock_get_secret:
mock_get_secret.return_value = None
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='in _complete_vertex_ai_beta') as exc_info:
completion(
model="gemini/gemini-1.5-flash",
messages=[{"role": "user", "content": "Test message"}],

View file

@ -109,7 +109,7 @@ class TestHostedVLLMRerankTransform:
)
assert url2 == "https://api.example.com/rerank"
# Raises if api_base is None
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='api_base must be provided for Hosted VLLM rerank'):
self.config.get_complete_url(None, self.model)
def test_transform_response(self):

View file

@ -46,7 +46,7 @@ def test_langflow_config_get_complete_url():
def test_langflow_config_get_complete_url_requires_api_base():
config = LangFlowConfig()
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='api_base is required for LangFlow\\. Set it via'):
config.get_complete_url(
api_base=None,
api_key=None,

View file

@ -154,7 +154,7 @@ class TestModelScopeImageGenerationTransformation:
mock_get_secret.return_value = None
headers = {}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='MODELSCOPE_API_KEY is not set\\. Please set it via') as exc_info:
self.config.validate_environment(
headers=headers,
model=self.model,
@ -367,7 +367,7 @@ class TestModelScopeImageGenerationTransformation:
model_response = ImageResponse(data=[])
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='litellm\\.BadRequestError: ModelScope error: Invalid prompt') as exc_info:
self.config.transform_image_generation_response(
model=self.model,
raw_response=mock_response,
@ -393,7 +393,7 @@ class TestModelScopeImageGenerationTransformation:
model_response = ImageResponse(data=[])
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='litellm\\.InternalServerError: Error parsing ModelScope') as exc_info:
self.config.transform_image_generation_response(
model=self.model,
raw_response=mock_response,

View file

@ -47,7 +47,7 @@ class TestNovitaConfig:
"""Test error handling when API key is missing"""
config = NovitaConfig()
with pytest.raises(ValueError) as excinfo:
with pytest.raises(ValueError, match='Missing Novita AI API Key - A call is being made to novita') as excinfo:
config.validate_environment(
headers={},
model="novita/meta-llama/llama-3.3-70b-instruct",

View file

@ -98,7 +98,7 @@ class TestOCIChatConfig:
config = OCIChatConfig()
headers = {}
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match='Missing required parameters: oci_user, oci_fingerprint') as excinfo:
config.validate_environment(
headers=headers,
model=TEST_MODEL,
@ -272,7 +272,7 @@ class TestOCIChatConfig:
"oci_serving_mode": "INVALID_MODE",
}
with pytest.raises(Exception) as excinfo:
with pytest.raises(Exception, match="kwarg `oci_serving_mode` must be either 'ON_DEMAND' or") as excinfo:
config.transform_request(
model=TEST_MODEL_NAME,
messages=TEST_MESSAGES, # type: ignore
@ -892,7 +892,7 @@ class TestOCISignerSupport:
optional_params = {"oci_signer": MockSigner(), "method": "INVALID"}
with pytest.raises(ValueError) as excinfo:
with pytest.raises(ValueError, match='Unsupported HTTP method: INVALID') as excinfo:
config.sign_request(
headers={},
optional_params=optional_params,
@ -1604,7 +1604,7 @@ class TestOCIKeyNormalization:
# We can't fully test signing without a real key, but we can verify
# the error message indicates the key was processed (not a type error)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='why-can-t-i-import-my-pem-file for more details\\.') as exc_info:
sign_with_manual_credentials(
headers={},
optional_params=optional_params,
@ -1630,7 +1630,7 @@ class TestOCIKeyNormalization:
"oci_key": crlf_pem,
}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='why-can-t-i-import-my-pem-file for more details\\.') as exc_info:
sign_with_manual_credentials(
headers={},
optional_params=optional_params,
@ -1692,7 +1692,7 @@ class TestOCIValidateEnvironment:
def test_missing_required_credentials_raises_error(self, config):
"""Test that missing required credentials raise an error."""
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Missing required parameters: oci_user, oci_fingerprint') as exc_info:
config.validate_environment(
headers={},
model="oci/xai.grok-3",
@ -1875,7 +1875,7 @@ class TestOCIImageUrlTransformation:
}
]
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Prop `image_url` must be a string or an object with a `url`') as exc_info:
adapt_messages_to_generic_oci_standard(messages)
assert "image_url" in str(exc_info.value)
@ -1899,7 +1899,7 @@ class TestOCIImageUrlTransformation:
}
]
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Prop `image_url` must be a string or an object with a `url`') as exc_info:
adapt_messages_to_generic_oci_standard(messages)
assert "image_url" in str(exc_info.value)

View file

@ -111,33 +111,51 @@ class TestCognitionProviderIdentity:
class TestCognitionCostTracking:
@pytest.mark.parametrize(
"model, input_cost, output_cost",
"model, input_cost, output_cost, cache_read_cost",
[
("cognition/swe-1.6", 5e-07, 2.5e-06),
("cognition/swe-1.7", 2.5e-06, 1.25e-05),
("cognition/swe-1.6", 5e-07, 2.5e-06, 2e-07),
("cognition/swe-1.7", 5e-07, 2.5e-06, 2e-07),
("cognition/swe-1.7-lightning", 2.5e-06, 1.25e-05, 1e-06),
],
)
def test_cost_map_entries(self, model: str, input_cost: float, output_cost: float):
def test_cost_map_entries(self, model: str, input_cost: float, output_cost: float, cache_read_cost: float):
info = litellm.get_model_info(model=model)
assert info["litellm_provider"] == "cognition"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == input_cost
assert info["output_cost_per_token"] == output_cost
assert info["cache_read_input_token_cost"] == cache_read_cost
def test_cost_differs_from_openai_pricing(self):
@pytest.mark.parametrize(
"model, expected_prompt_cost, expected_completion_cost",
[
("cognition/swe-1.7", 0.5, 2.5),
("cognition/swe-1.7-lightning", 2.5, 12.5),
],
)
def test_cost_differs_from_openai_pricing(
self, model: str, expected_prompt_cost: float, expected_completion_cost: float
):
"""A cognition-prefixed model must never be priced off an OpenAI cost entry."""
from litellm.cost_calculator import cost_per_token
prompt_cost, completion_cost = cost_per_token(
model="cognition/swe-1.7",
model=model,
prompt_tokens=1_000_000,
completion_tokens=1_000_000,
custom_llm_provider="cognition",
)
assert prompt_cost == pytest.approx(2.5)
assert completion_cost == pytest.approx(12.5)
assert prompt_cost == pytest.approx(expected_prompt_cost)
assert completion_cost == pytest.approx(expected_completion_cost)
def test_lightning_is_five_times_the_standard_tier(self):
standard = litellm.get_model_info(model="cognition/swe-1.7")
lightning = litellm.get_model_info(model="cognition/swe-1.7-lightning")
assert lightning["input_cost_per_token"] == pytest.approx(standard["input_cost_per_token"] * 5)
assert lightning["output_cost_per_token"] == pytest.approx(standard["output_cost_per_token"] * 5)
def test_supported_endpoints_matrix(self):
matrix = json.loads((Path(litellm.__file__).parent / "provider_endpoints_support_backup.json").read_text())
@ -170,6 +188,30 @@ class TestCognitionRouting:
mock_response="hello from swe",
)
usage = response.usage
expected = usage.prompt_tokens * 5e-07 + usage.completion_tokens * 2.5e-06
assert response._hidden_params["response_cost"] == pytest.approx(expected)
@pytest.mark.asyncio
async def test_router_spend_uses_the_lightning_entry_for_lightning(self):
"""The Lightning tier is its own model, costed off its own entry."""
from litellm import Router
router = Router(
model_list=[
{
"model_name": "swe-lightning",
"litellm_params": {"model": "cognition/swe-1.7-lightning", "api_key": "sk-test"},
}
]
)
response = await router.acompletion(
model="swe-lightning",
messages=[{"role": "user", "content": "hi"}],
mock_response="hello from swe lightning",
)
usage = response.usage
expected = usage.prompt_tokens * 2.5e-06 + usage.completion_tokens * 1.25e-05
assert response._hidden_params["response_cost"] == pytest.approx(expected)

View file

@ -42,7 +42,7 @@ class TestPGVectorStoreConfig:
litellm_params = GenericLiteLLMParams()
headers = {}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='PG Vector API key is required\\. Set PG_VECTOR_API_KEY') as exc_info:
config.validate_environment(headers, litellm_params)
assert "PG Vector API key is required" in str(exc_info.value)
@ -84,7 +84,7 @@ class TestPGVectorStoreConfig:
config = PGVectorStoreConfig()
litellm_params = {}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='PG Vector API base URL is required\\. Set') as exc_info:
config.get_complete_url(None, litellm_params)
assert "PG Vector API base URL is required" in str(exc_info.value)

View file

@ -167,7 +167,7 @@ class TestRecraftImageEditTransformation:
mock_response.status_code = 500
mock_response.headers = {}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Error transforming image edit response: Invalid JSON: line') as exc_info:
self.config.transform_image_edit_response(
model=self.model,
raw_response=mock_response,

View file

@ -64,7 +64,7 @@ class TestRecraftImageGenerationTransformation:
non_default_params = {"n": 2, "unsupported_param": "value"}
optional_params = {}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Supported parameters are') as exc_info:
self.config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
@ -171,7 +171,7 @@ class TestRecraftImageGenerationTransformation:
mock_get_secret.return_value = None
headers = {}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='RECRAFT_API_KEY is not set') as exc_info:
self.config.validate_environment(
headers=headers,
model=self.model,
@ -248,7 +248,7 @@ class TestRecraftImageGenerationTransformation:
model_response = ImageResponse(data=[])
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Error transforming image generation response: Invalid JSON') as exc_info:
self.config.transform_image_generation_response(
model=self.model,
raw_response=mock_response,

View file

@ -19,6 +19,8 @@ from unittest.mock import MagicMock
import httpx
import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from litellm.llms.sagemaker.chat.transformation import SagemakerChatConfig
@ -233,3 +235,85 @@ def test_decoder_reassembles_frames_across_arbitrary_byte_boundaries(split_size)
]
assert texts == [f"token{i} " for i in range(len(frames))]
_INFERENCE_COMPONENT_HEADER = "X-Amzn-SageMaker-Inference-Component"
_STUB_COMPLETION_RESPONSE = {
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 1700000000,
"model": "served-model",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
}
class _RequestCapturingHTTPHandler(HTTPHandler):
"""Injected transport that records exactly what sagemaker_chat put on the wire."""
def __init__(self) -> None:
super().__init__()
self.request_headers: dict[str, str] = {}
self.request_body: dict = {}
def post(self, url: str, headers=None, data=None, **kwargs) -> httpx.Response:
self.request_headers = dict(headers or {})
self.request_body = json.loads(data)
return httpx.Response(200, json=_STUB_COMPLETION_RESPONSE, request=httpx.Request("POST", url))
def _invoke_sagemaker_chat(monkeypatch, **extra_params) -> _RequestCapturingHTTPHandler:
"""Drive one sagemaker_chat completion against an injected transport.
A Bedrock API key short-circuits SigV4 inside `BaseAWSLLM._sign_request`, which would hide
whether the inference-component header is really covered by the signature, so it is cleared.
"""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
client = _RequestCapturingHTTPHandler()
litellm.completion(
model="sagemaker_chat/my-endpoint",
messages=[{"role": "user", "content": "hi"}],
aws_access_key_id="AKIATESTTESTTESTTEST",
aws_secret_access_key="test-secret-key",
aws_region_name="us-east-1",
client=client,
**extra_params,
)
return client
def test_model_id_is_sent_as_a_signed_inference_component_header(monkeypatch):
"""`model_id` names an inference component and must reach SageMaker as a signed header.
Endpoints backed by inference components reject any request without
`X-Amzn-SageMaker-Inference-Component` with HTTP 400 INFERENCE_COMPONENT_NAME_MISSING, so the
header has to be built before `sign_request` runs and end up inside SignedHeaders.
"""
client = _invoke_sagemaker_chat(monkeypatch, model_id="my-inference-component")
assert client.request_headers[_INFERENCE_COMPONENT_HEADER] == "my-inference-component"
assert "x-amzn-sagemaker-inference-component" in client.request_headers["Authorization"]
def test_no_inference_component_header_when_model_id_is_unset(monkeypatch):
"""Plain endpoints must not receive the header at all, not even an empty one."""
client = _invoke_sagemaker_chat(monkeypatch)
assert not any(name.lower() == _INFERENCE_COMPONENT_HEADER.lower() for name in client.request_headers)
def test_hf_model_name_becomes_the_body_model(monkeypatch):
"""`hf_model_name` names the served model, and containers that validate the body's `model`
404 on the endpoint name, so it has to replace it rather than ride along as an extra field."""
client = _invoke_sagemaker_chat(monkeypatch, hf_model_name="org/served-model")
assert client.request_body["model"] == "org/served-model"
assert "hf_model_name" not in client.request_body
def test_body_model_stays_the_endpoint_name_when_hf_model_name_is_unset(monkeypatch):
"""Without `hf_model_name` the body must keep the model it has today."""
client = _invoke_sagemaker_chat(monkeypatch)
assert client.request_body["model"] == "my-endpoint"

View file

@ -83,7 +83,7 @@ class TestStabilityImageGenerationConfig:
non_default_params = {"unsupported_param": "value"}
optional_params = {}
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match="Supported parameters are \\['n', 'size',") as exc_info:
self.config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
@ -168,7 +168,7 @@ class TestStabilityImageGenerationConfig:
def test_validate_environment_raises_without_api_key(self):
"""Test that validate_environment raises error without API key"""
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='STABILITY_API_KEY is not set\\. Please set it via') as exc_info:
self.config.validate_environment(
headers={},
model="stability/sd3",
@ -251,7 +251,7 @@ class TestStabilityImageGenerationConfig:
model_response = ImageResponse(data=[])
mock_logging = MagicMock()
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Content was filtered by Stability AI safety systems') as exc_info:
self.config.transform_image_generation_response(
model="stability/sd3",
raw_response=mock_response,

View file

@ -697,7 +697,7 @@ class TestErrorHandling:
}
}
mock_response = _make_mock_response(body, status_code=400)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='TinyFish Search: query is required\\. See https') as exc_info:
config.transform_search_response(
raw_response=mock_response, logging_obj=None
)
@ -713,7 +713,7 @@ class TestErrorHandling:
mock_response = _make_mock_response(
body, status_code=429, headers={"Retry-After": "60"}
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='TinyFish Search: rate limit exceeded\\. See https') as exc_info:
config.transform_search_response(
raw_response=mock_response, logging_obj=None
)
@ -728,7 +728,7 @@ class TestErrorHandling:
config = TinyfishSearchConfig()
body = {"errors": [{"code": "10000", "message": "Internal"}]}
mock_response = _make_mock_response(body, status_code=502)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='TinyFish Search') as exc_info:
config.transform_search_response(
raw_response=mock_response, logging_obj=None
)
@ -742,7 +742,7 @@ class TestErrorHandling:
mock_response = _make_mock_response(
json_data=None, status_code=502, text="<html>Bad Gateway</html>"
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='TinyFish Search: <html>Bad Gateway<') as exc_info:
config.transform_search_response(
raw_response=mock_response, logging_obj=None
)
@ -756,7 +756,7 @@ class TestErrorHandling:
mock_response = _make_mock_response(
json_data=None, status_code=200, text="not json"
)
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='TinyFish Search: Expected JSON response, got: not json\\.') as exc_info:
config.transform_search_response(
raw_response=mock_response, logging_obj=None
)
@ -785,7 +785,7 @@ class TestErrorHandling:
# check TinyFish's schema, not their own input.
config = TinyfishSearchConfig()
mock_response = _make_mock_response({"query": "x"}) # no `results` key
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='validation error for SearchResponse') as exc_info:
config.transform_search_response(
raw_response=mock_response, logging_obj=None
)

View file

@ -159,7 +159,7 @@ class TestVertexAIFilesIntegration:
# This test ensures the type annotations and error messages include vertex_ai
# Test that calling with unsupported provider raises appropriate error
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match="unsupported_provider' is not a valid LlmProviders") as exc_info:
litellm.file_content(
file_id="test-file-id",
custom_llm_provider="unsupported_provider", # This should fail

View file

@ -33,7 +33,7 @@ def test_validate_vertex_location_accepts_valid(location):
["attacker.example/", "evil.com#", "us.attacker.example", "us/../..", "US", "us_central1", "-us", "", None],
)
def test_validate_vertex_location_rejects_invalid(location):
with pytest.raises(ValueError):
with pytest.raises(ValueError, match="vertex_location is required|Invalid vertex_location format"):
validate_vertex_location(location)

View file

@ -137,7 +137,7 @@ class TestVolcengineResponsesAPITransformation:
monkeypatch.delenv("ARK_API_KEY", raising=False)
monkeypatch.delenv("VOLCENGINE_API_KEY", raising=False)
with pytest.raises(ValueError):
with pytest.raises(ValueError, match='Volcengine API key is required\\. Set ARK_API_KEY /'):
config.validate_environment(headers={}, model="volcengine/demo", litellm_params={})
def test_unsupported_params_are_dropped_with_extra_body(self):

View file

@ -202,7 +202,7 @@ def test_volcengine_embedding_error_scenarios():
k: v for k, v in scenario.items() if k != "expected_error_pattern"
}
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match=f"(?i){scenario['expected_error_pattern']}") as exc_info:
litellm.embedding(input=["test"], **test_params)
# Verify error message contains expected pattern

View file

@ -227,7 +227,7 @@ class TestVoyageRerankTransform:
mock_logging = MagicMock()
model_response = RerankResponse()
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Unauthorized') as exc_info:
self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,
@ -248,7 +248,7 @@ class TestVoyageRerankTransform:
mock_logging = MagicMock()
model_response = RerankResponse()
with pytest.raises(Exception) as exc_info:
with pytest.raises(Exception, match='Failed to parse response: Invalid JSON response') as exc_info:
self.config.transform_rerank_response(
model=self.model,
raw_response=mock_response,

View file

@ -195,7 +195,7 @@ class TestVoyageMultimodalEmbeddings:
monkeypatch.setattr(module, "get_secret_str", lambda name: None)
config = VoyageMultimodalEmbeddingConfig()
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Voyage API key is required for multimodal embeddings\\. Set') as exc_info:
config.validate_environment(
{}, "voyage-multimodal-3.5", [], {}, {}, api_key=None
)
@ -207,7 +207,7 @@ class TestVoyageMultimodalEmbeddings:
)
config = VoyageMultimodalEmbeddingConfig()
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='Voyage multimodal embeddings require a non-empty') as exc_info:
config._normalize_content_item({"type": "image_url", "image_url": {}})
assert "image_url" in str(exc_info.value)

View file

@ -168,7 +168,7 @@ def test_responses_config_raises_when_no_key_is_available(monkeypatch):
monkeypatch.setattr(litellm, "api_key", None)
monkeypatch.delenv("XAI_API_KEY", raising=False)
with pytest.raises(ValueError) as exc_info:
with pytest.raises(ValueError, match='XAI API key is required\\. Set api_key, litellm\\.xai_key') as exc_info:
XAIResponsesAPIConfig().validate_environment({}, "xai/grok-3-mini", None)
error_message = str(exc_info.value)

Some files were not shown because too many files have changed in this diff Show more