mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
commit
1cef8823fa
165 changed files with 2002 additions and 1018 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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={}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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={})
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"})
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"}],
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue