Merge branch 'litellm_contributor_prs_09_18_2025_p2' into fix/issue-14685-bedrock-titan-v2-encoding-format

This commit is contained in:
Krish Dholakia 2025-09-18 17:54:33 -07:00 • committed by GitHub
commit 63c26d7a4f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
27 changed files with 876 additions and 363 deletions

View file

@ -1458,6 +1458,7 @@ jobs:
# - run: python ./tests/documentation_tests/test_general_setting_keys.py
- run: python ./tests/code_coverage_tests/check_licenses.py
- run: python ./tests/code_coverage_tests/router_code_coverage.py
- run: python ./tests/code_coverage_tests/test_chat_completion_imports.py
- run: python ./tests/code_coverage_tests/info_log_check.py
- run: python ./tests/code_coverage_tests/test_ban_set_verbose.py
- run: python ./tests/code_coverage_tests/code_qa_check_tests.py

View file

@ -1821,6 +1821,7 @@ Here's an example of using a bedrock model with LiteLLM. For a complete list, re
| Mistral 7B Instruct | `completion(model='bedrock/mistral.mistral-7b-instruct-v0:2', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
| Mixtral 8x7B Instruct | `completion(model='bedrock/mistral.mixtral-8x7b-instruct-v0:1', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` |
## Bedrock Embedding
### API keys

View file

@ -0,0 +1,95 @@
## Bedrock Embedding
## Supported Embedding Models
| Provider | LiteLLM Route | AWS Documentation |
|----------|---------------|-------------------|
| Amazon Titan | `bedrock/amazon.*` | [Amazon Titan Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/titan-embedding-models.html) |
| Cohere | `bedrock/cohere.*` | [Cohere Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-embed.html) |
| TwelveLabs | `bedrock/us.twelvelabs.*` | [TwelveLabs](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-twelvelabs.html) |
### API keys
This can be set as env variables or passed as **params to litellm.embedding()**
```python
import os
os.environ["AWS_ACCESS_KEY_ID"] = "" # Access key
os.environ["AWS_SECRET_ACCESS_KEY"] = "" # Secret access key
os.environ["AWS_REGION_NAME"] = "" # us-east-1, us-east-2, us-west-1, us-west-2
```
## Usage
### LiteLLM Python SDK
```python
from litellm import embedding
response = embedding(
model="bedrock/amazon.titan-embed-text-v1",
input=["good morning from litellm"],
)
print(response)
```
### LiteLLM Proxy Server
#### 1. Setup config.yaml
```yaml
model_list:
- model_name: titan-embed-v1
litellm_params:
model: bedrock/amazon.titan-embed-text-v1
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: titan-embed-v2
litellm_params:
model: bedrock/amazon.titan-embed-text-v2:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
```
#### 2. Start Proxy
```bash
litellm --config /path/to/config.yaml
```
#### 3. Use with OpenAI Python SDK
```python
import openai
client = openai.OpenAI(
api_key="anything",
base_url="http://0.0.0.0:4000"
)
response = client.embeddings.create(
input=["good morning from litellm"],
model="titan-embed-v1"
)
print(response)
```
#### 4. Use with LiteLLM Python SDK
```python
import litellm
response = litellm.embedding(
model="titan-embed-v1", # model alias from config.yaml
input=["good morning from litellm"],
api_base="http://0.0.0.0:4000",
api_key="anything"
)
print(response)
```
## Supported AWS Bedrock Embedding Models
| Model Name | Usage | Supported Additional OpenAI params |
|----------------------|---------------------------------------------|-----|
| Titan Embeddings V2 | `embedding(model="bedrock/amazon.titan-embed-text-v2:0", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_v2_transformation.py#L59) |
| Titan Embeddings - V1 | `embedding(model="bedrock/amazon.titan-embed-text-v1", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_g1_transformation.py#L53)
| Titan Multimodal Embeddings | `embedding(model="bedrock/amazon.titan-embed-image-v1", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/amazon_titan_multimodal_transformation.py#L28) |
| TwelveLabs Marengo Embed 2.7 | `embedding(model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0", input=input)` | Supports multimodal input (text, video, audio, image) |
| Cohere Embeddings - English | `embedding(model="bedrock/cohere.embed-english-v3", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/cohere_transformation.py#L18)
| Cohere Embeddings - Multilingual | `embedding(model="bedrock/cohere.embed-multilingual-v3", input=input)` | [here](https://github.com/BerriAI/litellm/blob/f5905e100068e7a4d61441d7453d7cf5609c2121/litellm/llms/bedrock/embed/cohere_transformation.py#L18)
### Advanced - [Drop Unsupported Params](https://docs.litellm.ai/docs/completion/drop_params#openai-proxy-usage)
### Advanced - [Pass model/provider-specific Params](https://docs.litellm.ai/docs/completion/provider_specific_params#proxy-usage)

View file

@ -411,6 +411,7 @@ const sidebars = {
label: "Bedrock",
items: [
"providers/bedrock",
"providers/bedrock_embedding",
"providers/bedrock_agents",
"providers/bedrock_batches",
"providers/bedrock_vector_store",

View file

@ -67,6 +67,7 @@ from litellm.constants import (
bedrock_embedding_models,
known_tokenizer_config,
BEDROCK_INVOKE_PROVIDERS_LITERAL,
BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
BEDROCK_CONVERSE_MODELS,
DEFAULT_MAX_TOKENS,
DEFAULT_SOFT_BUDGET,

View file

@ -769,6 +769,12 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[
"deepseek_r1",
]
BEDROCK_EMBEDDING_PROVIDERS_LITERAL = Literal[
"cohere",
"amazon",
"twelvelabs",
]
BEDROCK_CONVERSE_MODELS = [
"openai.gpt-oss-20b-1:0",
"openai.gpt-oss-120b-1:0",
@ -822,6 +828,7 @@ bedrock_embedding_models: set = set(
"amazon.titan-embed-text-v1",
"cohere.embed-english-v3",
"cohere.embed-multilingual-v3",
"twelvelabs.marengo-embed-2-7-v1:0",
]
)
@ -1065,4 +1072,6 @@ SENTRY_PII_DENYLIST = [
]
# CoroutineChecker cache configuration
COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY = int(os.getenv("COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY", 1000))
COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY = int(
os.getenv("COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY", 1000)
)

View file

@ -498,6 +498,7 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
"guardrail_information": standard_logging_payload.get(
"guardrail_information", None
),
"is_streamed_request": self._get_stream_value_from_payload(standard_logging_payload),
}
#########################################################
@ -561,6 +562,31 @@ class DataDogLLMObsLogger(DataDogLogger, CustomBatchLogger):
return latency_metrics
def _get_stream_value_from_payload(self, standard_logging_payload: StandardLoggingPayload) -> bool:
"""
Extract the stream value from standard logging payload.
The stream field in StandardLoggingPayload is only set to True for completed streaming responses.
For non-streaming requests, it's None. The original stream parameter is in model_parameters.
Returns:
bool: True if this was a streaming request, False otherwise
"""
# Check top-level stream field first (only True for completed streaming)
stream_value = standard_logging_payload.get("stream")
if stream_value is True:
return True
# Fallback to model_parameters.stream for original request parameters
model_params = standard_logging_payload.get("model_parameters", {})
if isinstance(model_params, dict):
stream_value = model_params.get("stream")
if stream_value is True:
return True
# Default to False for non-streaming requests
return False
def _get_spend_metrics(
self, standard_logging_payload: StandardLoggingPayload
) -> DDLLMObsSpendMetrics:

View file

@ -0,0 +1,56 @@
"""
Cached imports module for LiteLLM.
This module provides cached import functionality to avoid repeated imports
inside functions that are critical to performance.
"""
from typing import TYPE_CHECKING, Callable, Optional, Type
# Type annotations for cached imports
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.coroutine_checker import CoroutineChecker
# Global cache variables
_LiteLLMLogging: Optional[Type["Logging"]] = None
_coroutine_checker: Optional["CoroutineChecker"] = None
_set_callbacks: Optional[Callable] = None
def get_litellm_logging_class() -> Type["Logging"]:
"""Get the cached LiteLLM Logging class, initializing if needed."""
global _LiteLLMLogging
if _LiteLLMLogging is not None:
return _LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import Logging
_LiteLLMLogging = Logging
return _LiteLLMLogging
def get_coroutine_checker() -> "CoroutineChecker":
"""Get the cached coroutine checker instance, initializing if needed."""
global _coroutine_checker
if _coroutine_checker is not None:
return _coroutine_checker
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
_coroutine_checker = coroutine_checker
return _coroutine_checker
def get_set_callbacks() -> Callable:
"""Get the cached set_callbacks function, initializing if needed."""
global _set_callbacks
if _set_callbacks is not None:
return _set_callbacks
from litellm.litellm_core_utils.litellm_logging import set_callbacks
_set_callbacks = set_callbacks
return _set_callbacks
def clear_cached_imports() -> None:
"""Clear all cached imports. Useful for testing or memory management."""
global _LiteLLMLogging, _coroutine_checker, _set_callbacks
_LiteLLMLogging = None
_coroutine_checker = None
_set_callbacks = None

View file

@ -556,7 +556,7 @@ def exception_type( # type: ignore # noqa: PLR0915
model=model,
llm_provider="anthropic",
)
elif "overloaded_error" in error_str:
elif "overloaded_error" in error_str or "Overloaded" in error_str:
exception_mapping_worked = True
raise InternalServerError(
message="AnthropicError - {}".format(error_str),
@ -1449,6 +1449,14 @@ def exception_type( # type: ignore # noqa: PLR0915
model=model,
response=getattr(original_exception, "response", None),
)
elif "invalid type: parameter" in error_str:
exception_mapping_worked = True
raise BadRequestError(
message=f"CohereException - {original_exception.message}",
llm_provider="cohere",
model=model,
response=getattr(original_exception, "response", None),
)
elif "too many tokens" in error_str:
exception_mapping_worked = True
raise ContextWindowExceededError(

View file

@ -3079,7 +3079,6 @@ class BedrockConverseMessagesProcessor:
messages.append(DEFAULT_USER_CONTINUE_MESSAGE)
return messages
@staticmethod
async def _bedrock_converse_messages_pt_async( # noqa: PLR0915
messages: List,
@ -3124,9 +3123,9 @@ class BedrockConverseMessagesProcessor:
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "guarded_text":
# Wrap guarded_text in guardrailConverseContent block
# Wrap guarded_text in guardContent block
_part = BedrockContentBlock(
guardrailConverseContent={"text": element["text"]}
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] == "image_url":
@ -3171,7 +3170,6 @@ class BedrockConverseMessagesProcessor:
msg_i += 1
if user_content:
if len(contents) > 0 and contents[-1]["role"] == "user":
if (
assistant_continue_message is not None
@ -3506,9 +3504,9 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
_part = BedrockContentBlock(text=element["text"])
_parts.append(_part)
elif element["type"] == "guarded_text":
# Wrap guarded_text in guardrailConverseContent block
# Wrap guarded_text in guardContent block
_part = BedrockContentBlock(
guardrailConverseContent={"text": element["text"]}
guardContent={"text": {"text": element["text"]}}
)
_parts.append(_part)
elif element["type"] == "image_url":
@ -3554,7 +3552,6 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
msg_i += 1
if user_content:
if len(contents) > 0 and contents[-1]["role"] == "user":
if (
assistant_continue_message is not None

View file

@ -20,7 +20,11 @@ from pydantic import BaseModel
from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.constants import BEDROCK_INVOKE_PROVIDERS_LITERAL, BEDROCK_MAX_POLICY_SIZE
from litellm.constants import (
BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
BEDROCK_INVOKE_PROVIDERS_LITERAL,
BEDROCK_MAX_POLICY_SIZE,
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.secret_managers.main import get_secret, get_secret_str
@ -327,6 +331,40 @@ class BaseAWSLLM:
return provider
return None
@staticmethod
def get_bedrock_embedding_provider(
model: str,
) -> Optional[BEDROCK_EMBEDDING_PROVIDERS_LITERAL]:
"""
Helper function to get the bedrock embedding provider from the model
Handles scenarios like:
1. model=cohere.embed-english-v3:0 -> Returns `cohere`
2. model=amazon.titan-embed-text-v1 -> Returns `amazon`
3. model=us.twelvelabs.marengo-embed-2-7-v1:0 -> Returns `twelvelabs`
4. model=twelvelabs.marengo-embed-2-7-v1:0 -> Returns `twelvelabs`
"""
# Handle regional models like us.twelvelabs.marengo-embed-2-7-v1:0
if "." in model:
parts = model.split(".")
# Check if the second part (after potential region) is a known provider
if len(parts) >= 2:
potential_provider = parts[1] # e.g., "twelvelabs" from "us.twelvelabs.marengo-embed-2-7-v1:0"
if potential_provider in get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL):
return cast(BEDROCK_EMBEDDING_PROVIDERS_LITERAL, potential_provider)
# Check if the first part is a known provider (standard format)
potential_provider = parts[0] # e.g., "cohere" from "cohere.embed-english-v3:0"
if potential_provider in get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL):
return cast(BEDROCK_EMBEDDING_PROVIDERS_LITERAL, potential_provider)
# Fallback: check if any provider name appears in the model string
for provider in get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL):
if provider in model:
return cast(BEDROCK_EMBEDDING_PROVIDERS_LITERAL, provider)
return None
def _get_aws_region_name(
self,
optional_params: dict,

View file

@ -4,12 +4,13 @@ Handles embedding calls to Bedrock's `/invoke` endpoint
import copy
import json
from typing import Any, Callable, List, Optional, Tuple, Union
import urllib.parse
from typing import Any, Callable, List, Optional, Tuple, Union, get_args
import httpx
import litellm
from litellm.constants import BEDROCK_EMBEDDING_PROVIDERS_LITERAL
from litellm.llms.cohere.embed.handler import embedding as cohere_embedding
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -18,7 +19,11 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
)
from litellm.secret_managers.main import get_secret
from litellm.types.llms.bedrock import AmazonEmbeddingRequest, CohereEmbeddingRequest
from litellm.types.llms.bedrock import (
AmazonEmbeddingRequest,
CohereEmbeddingRequest,
TwelveLabsMarengoEmbeddingRequest,
)
from litellm.types.utils import EmbeddingResponse
from ..base_aws_llm import BaseAWSLLM
@ -29,6 +34,7 @@ from .amazon_titan_multimodal_transformation import (
)
from .amazon_titan_v2_transformation import AmazonTitanV2Config
from .cohere_transformation import BedrockCohereEmbeddingConfig
from .twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig
class BedrockEmbedding(BaseAWSLLM):
@ -145,6 +151,44 @@ class BedrockEmbedding(BaseAWSLLM):
raise BedrockError(status_code=408, message="Timeout error occurred.")
return response.json()
def _transform_response(
self, response_list: List[dict], model: str, provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL
) -> Optional[EmbeddingResponse]:
"""
Transforms the response from the Bedrock embedding provider to the OpenAI format.
"""
returned_response: Optional[EmbeddingResponse] = None
if model == "amazon.titan-embed-image-v1":
returned_response = (
AmazonTitanMultimodalEmbeddingG1Config()._transform_response(
response_list=response_list, model=model
)
)
elif model == "amazon.titan-embed-text-v1":
returned_response = AmazonTitanG1Config()._transform_response(
response_list=response_list, model=model
)
elif model == "amazon.titan-embed-text-v2:0":
returned_response = AmazonTitanV2Config()._transform_response(
response_list=response_list, model=model
)
elif provider == "twelvelabs":
returned_response = TwelveLabsMarengoEmbeddingConfig()._transform_response(
response_list=response_list, model=model
)
##########################################################
# Validate returned response
##########################################################
if returned_response is None:
raise Exception(
"Unable to map model response to known provider format. model={}".format(
model
)
)
return returned_response
def _single_func_embeddings(
self,
@ -157,6 +201,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_region_name: str,
model: str,
logging_obj: Any,
provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
api_key: Optional[str] = None,
):
responses: List[dict] = []
@ -164,16 +209,16 @@ class BedrockEmbedding(BaseAWSLLM):
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=json.dumps(data),
headers=headers,
api_key=api_key
)
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=json.dumps(data),
headers=headers,
api_key=api_key,
)
## LOGGING
logging_obj.pre_call(
@ -203,32 +248,9 @@ class BedrockEmbedding(BaseAWSLLM):
responses.append(response)
returned_response: Optional[EmbeddingResponse] = None
## TRANSFORM RESPONSE ##
if model == "amazon.titan-embed-image-v1":
returned_response = (
AmazonTitanMultimodalEmbeddingG1Config()._transform_response(
response_list=responses, model=model
)
)
elif model == "amazon.titan-embed-text-v1":
returned_response = AmazonTitanG1Config()._transform_response(
response_list=responses, model=model
)
elif model == "amazon.titan-embed-text-v2:0":
returned_response = AmazonTitanV2Config()._transform_response(
response_list=responses, model=model
)
if returned_response is None:
raise Exception(
"Unable to map model response to known provider format. model={}".format(
model
)
)
return returned_response
return self._transform_response(
response_list=responses, model=model, provider=provider
)
async def _async_single_func_embeddings(
self,
@ -241,6 +263,7 @@ class BedrockEmbedding(BaseAWSLLM):
aws_region_name: str,
model: str,
logging_obj: Any,
provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
api_key: Optional[str] = None,
):
responses: List[dict] = []
@ -248,16 +271,16 @@ class BedrockEmbedding(BaseAWSLLM):
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=json.dumps(data),
headers=headers,
api_key=api_key,
)
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=json.dumps(data),
headers=headers,
api_key=api_key,
)
## LOGGING
logging_obj.pre_call(
@ -286,33 +309,10 @@ class BedrockEmbedding(BaseAWSLLM):
)
responses.append(response)
returned_response: Optional[EmbeddingResponse] = None
## TRANSFORM RESPONSE ##
if model == "amazon.titan-embed-image-v1":
returned_response = (
AmazonTitanMultimodalEmbeddingG1Config()._transform_response(
response_list=responses, model=model
)
)
elif model == "amazon.titan-embed-text-v1":
returned_response = AmazonTitanG1Config()._transform_response(
response_list=responses, model=model
)
elif model == "amazon.titan-embed-text-v2:0":
returned_response = AmazonTitanV2Config()._transform_response(
response_list=responses, model=model
)
if returned_response is None:
raise Exception(
"Unable to map model response to known provider format. model={}".format(
model
)
)
return returned_response
return self._transform_response(
response_list=responses, model=model, provider=provider
)
def embeddings(
self,
@ -336,7 +336,7 @@ class BedrockEmbedding(BaseAWSLLM):
### TRANSFORMATION ###
unencoded_model_id = (
optional_params.pop("model_id", None) or model
) # default to model if not passed
) # default to model if not passed
modelId = urllib.parse.quote(unencoded_model_id, safe="")
aws_region_name = self._get_aws_region_name(
optional_params=optional_params,
@ -344,7 +344,12 @@ class BedrockEmbedding(BaseAWSLLM):
model_id=unencoded_model_id,
)
provider = model.split(".")[0]
provider = self.get_bedrock_embedding_provider(model)
if provider is None:
raise Exception(
f"Unable to determine bedrock embedding provider for model: {model}. "
f"Supported providers: {list(get_args(BEDROCK_EMBEDDING_PROVIDERS_LITERAL))}"
)
inference_params = copy.deepcopy(optional_params)
inference_params = {
k: v
@ -394,6 +399,15 @@ class BedrockEmbedding(BaseAWSLLM):
)
)
batch_data.append(transformed_request)
elif provider == "twelvelabs":
batch_data = []
for i in input:
twelvelabs_request: (
TwelveLabsMarengoEmbeddingRequest
) = TwelveLabsMarengoEmbeddingConfig()._transform_request(
input=i, inference_params=inference_params
)
batch_data.append(twelvelabs_request)
### SET RUNTIME ENDPOINT ###
endpoint_url, proxy_endpoint_url = self.get_runtime_endpoint(
@ -422,8 +436,9 @@ class BedrockEmbedding(BaseAWSLLM):
model=model,
logging_obj=logging_obj,
api_key=api_key,
provider=provider,
)
return self._single_func_embeddings(
returned_response = self._single_func_embeddings(
client=(
client
if client is not None and isinstance(client, HTTPHandler)
@ -438,14 +453,18 @@ class BedrockEmbedding(BaseAWSLLM):
model=model,
logging_obj=logging_obj,
api_key=api_key,
provider=provider,
)
if returned_response is None:
raise Exception("Unable to map Bedrock request to provider")
return returned_response
elif data is None:
raise Exception("Unable to map Bedrock request to provider")
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,

View file

@ -0,0 +1,140 @@
"""
Transformation logic from OpenAI /v1/embeddings format to Bedrock TwelveLabs Marengo /invoke format.
Why separate file? Make it easy to see how transformation works
Docs - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo.html
"""
from typing import List
from litellm.types.llms.bedrock import (
TwelveLabsMarengoEmbeddingRequest,
)
from litellm.types.utils import Embedding, EmbeddingResponse, Usage
from litellm.utils import get_base64_str, is_base64_encoded
class TwelveLabsMarengoEmbeddingConfig:
"""
Reference - https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-marengo.html
Supports text and image inputs for Phase 1.
Video and audio support will be added in Phase 2.
"""
def __init__(self) -> None:
pass
def get_supported_openai_params(self) -> List[str]:
return ["encoding_format", "textTruncate", "embeddingOption"]
def map_openai_params(
self, non_default_params: dict, optional_params: dict
) -> dict:
for k, v in non_default_params.items():
if k == "encoding_format":
# TwelveLabs doesn't have encoding_format, but we can map it to embeddingOption
if v == "float":
optional_params["embeddingOption"] = ["visual-text", "visual-image"]
elif k == "textTruncate":
optional_params["textTruncate"] = v
elif k == "embeddingOption":
optional_params["embeddingOption"] = v
return optional_params
def _transform_request(
self, input: str, inference_params: dict
) -> TwelveLabsMarengoEmbeddingRequest:
"""
Transform OpenAI-style input to TwelveLabs Marengo format.
Phase 1: Supports text and image inputs only.
"""
# Check if input is base64 encoded image
is_encoded = is_base64_encoded(input)
if is_encoded:
# Image input
b64_str = get_base64_str(input)
transformed_request = TwelveLabsMarengoEmbeddingRequest(
inputType="image", mediaSource={"base64String": b64_str}
)
else:
# Text input
transformed_request = TwelveLabsMarengoEmbeddingRequest(
inputType="text", inputText=input
)
# Set default textTruncate if not specified
if "textTruncate" not in inference_params:
transformed_request["textTruncate"] = "end"
# Apply any additional inference parameters
for k, v in inference_params.items():
if k not in [
"inputType",
"inputText",
"mediaSource",
]: # Don't override core fields
transformed_request[k] = v # type: ignore
return transformed_request
def _transform_response(
self, response_list: List[dict], model: str
) -> EmbeddingResponse:
"""
Transform TwelveLabs response to OpenAI format.
Handles the actual TwelveLabs response format: {"data": [{"embedding": [...]}]}
"""
embeddings: List[Embedding] = []
total_tokens = 0
for response in response_list:
# TwelveLabs response format has a "data" field containing the embeddings
if "data" in response and isinstance(response["data"], list):
for item in response["data"]:
if "embedding" in item:
# Single embedding response
embedding = Embedding(
embedding=item["embedding"],
index=len(embeddings),
object="embedding",
)
embeddings.append(embedding)
# Estimate token count (rough approximation)
if "inputTextTokenCount" in item:
total_tokens += item["inputTextTokenCount"]
else:
# Rough estimate: 1 token per 4 characters for text, or use embedding size
total_tokens += len(item["embedding"]) // 4
elif "embedding" in response:
# Direct embedding response (fallback for other formats)
embedding = Embedding(
embedding=response["embedding"],
index=len(embeddings),
object="embedding",
)
embeddings.append(embedding)
# Estimate token count (rough approximation)
if "inputTextTokenCount" in response:
total_tokens += response["inputTextTokenCount"]
else:
# Rough estimate: 1 token per 4 characters for text
total_tokens += len(response.get("inputText", "")) // 4
elif "embeddings" in response:
# Multiple embeddings response (from video/audio)
for i, emb in enumerate(response["embeddings"]):
embedding = Embedding(
embedding=emb["embedding"],
index=len(embeddings),
object="embedding",
)
embeddings.append(embedding)
total_tokens += len(emb["embedding"]) // 4 # Rough estimate
usage = Usage(prompt_tokens=total_tokens, total_tokens=total_tokens)
return EmbeddingResponse(data=embeddings, model=model, usage=usage)

View file

@ -296,6 +296,66 @@
"output_cost_per_token": 0.0,
"output_vector_size": 1024
},
"twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"us.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"eu.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"us.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"eu.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"amazon.titan-text-express-v1": {
"input_cost_per_token": 1.3e-06,
"litellm_provider": "bedrock",

View file

@ -1,4 +1,5 @@
import json
import re
from typing import Any, Dict, List, Optional
import orjson
@ -51,8 +52,6 @@ async def _read_request_body(request: Optional[Request]) -> Dict:
body_str = body.decode("utf-8") if isinstance(body, bytes) else body
# Replace invalid surrogate pairs
import re
# This regex finds incomplete surrogate pairs
body_str = re.sub(
r"[\uD800-\uDBFF](?![\uDC00-\uDFFF])", "", body_str

View file

@ -88,10 +88,14 @@ class BedrockConverseReasoningContentBlockDelta(TypedDict, total=False):
text: str
class GuardrailConverseTextBlock(TypedDict, total=False):
text: str
class GuardrailConverseContentBlock(TypedDict, total=False):
"""Content block for selective guardrail evaluation in Bedrock Converse API"""
text: str
text: GuardrailConverseTextBlock
class ContentBlock(TypedDict, total=False):
@ -103,7 +107,7 @@ class ContentBlock(TypedDict, total=False):
toolUse: ToolUseBlock
cachePoint: CachePointBlock
reasoningContent: BedrockConverseReasoningContentBlock
guardrailConverseContent: GuardrailConverseContentBlock
guardContent: GuardrailConverseContentBlock
class MessageBlock(TypedDict):
@ -367,6 +371,35 @@ class AmazonTitanMultimodalEmbeddingResponse(TypedDict):
message: str # Specifies any errors that occur during generation.
# TwelveLabs Marengo Embed 2.7 types
TWELVELABS_EMBEDDING_INPUT_TYPES = Literal["text", "image", "video", "audio"]
TWELVELABS_EMBEDDING_OPTIONS = Literal["visual-text", "visual-image", "audio"]
class TwelveLabsMediaSource(TypedDict, total=False):
base64String: str
s3Location: dict # {"uri": str, "bucketOwner": str}
class TwelveLabsMarengoEmbeddingRequest(TypedDict, total=False):
inputType: Required[TWELVELABS_EMBEDDING_INPUT_TYPES]
inputText: str
mediaSource: TwelveLabsMediaSource
textTruncate: Literal["end", "none"]
startSec: float
lengthSec: float
useFixedLengthSec: float
minClipSec: int
embeddingOption: List[TWELVELABS_EMBEDDING_OPTIONS]
class TwelveLabsMarengoEmbeddingResponse(TypedDict):
embedding: List[float]
embeddingOption: TWELVELABS_EMBEDDING_OPTIONS
startSec: float
endSec: float
AmazonEmbeddingRequest = Union[
AmazonTitanMultimodalEmbeddingRequest,
AmazonTitanV2EmbeddingRequest,

View file

@ -59,6 +59,12 @@ import litellm.litellm_core_utils.audio_utils.utils
import litellm.litellm_core_utils.json_validation_rule
import litellm.llms
import litellm.llms.gemini
# Import cached imports utilities
from litellm.litellm_core_utils.cached_imports import (
get_coroutine_checker,
get_litellm_logging_class,
get_set_callbacks,
)
from litellm.caching._internal_lru_cache import lru_cache_wrapper
from litellm.caching.caching import DualCache
from litellm.caching.caching_handler import CachingHandlerResponse, LLMCachingHandler
@ -222,6 +228,7 @@ from typing import (
get_args,
)
from openai import OpenAIError as OriginalError
from litellm.litellm_core_utils.thread_pool_executor import executor
@ -521,16 +528,12 @@ def get_dynamic_callbacks(
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
def function_setup( # noqa: PLR0915
original_function: str, rules_obj, start_time, *args, **kwargs
): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc.
### NOTICES ###
from litellm import Logging as LiteLLMLogging
from litellm.litellm_core_utils.litellm_logging import set_callbacks
if litellm.set_verbose is True:
verbose_logger.warning(
"`litellm.set_verbose` is deprecated. Please set `os.environ['LITELLM_LOG'] = 'DEBUG'` for debug logs."
@ -593,12 +596,12 @@ def function_setup( # noqa: PLR0915
+ litellm.failure_callback
)
)
set_callbacks(callback_list=callback_list, function_id=function_id)
get_set_callbacks()(callback_list=callback_list, function_id=function_id)
## ASYNC CALLBACKS
if len(litellm.input_callback) > 0:
removed_async_items = []
for index, callback in enumerate(litellm.input_callback): # type: ignore
if coroutine_checker.is_async_callable(callback):
if get_coroutine_checker().is_async_callable(callback):
litellm._async_input_callback.append(callback)
removed_async_items.append(index)
@ -608,7 +611,7 @@ def function_setup( # noqa: PLR0915
if len(litellm.success_callback) > 0:
removed_async_items = []
for index, callback in enumerate(litellm.success_callback): # type: ignore
if coroutine_checker.is_async_callable(callback):
if get_coroutine_checker().is_async_callable(callback):
litellm.logging_callback_manager.add_litellm_async_success_callback(
callback
)
@ -633,7 +636,7 @@ def function_setup( # noqa: PLR0915
if len(litellm.failure_callback) > 0:
removed_async_items = []
for index, callback in enumerate(litellm.failure_callback): # type: ignore
if coroutine_checker.is_async_callable(callback):
if get_coroutine_checker().is_async_callable(callback):
litellm.logging_callback_manager.add_litellm_async_failure_callback(
callback
)
@ -666,7 +669,7 @@ def function_setup( # noqa: PLR0915
removed_async_items = []
for index, callback in enumerate(kwargs["success_callback"]):
if (
coroutine_checker.is_async_callable(callback)
get_coroutine_checker().is_async_callable(callback)
or callback == "dynamodb"
or callback == "s3"
):
@ -790,7 +793,7 @@ def function_setup( # noqa: PLR0915
call_type=call_type,
):
stream = True
logging_obj = LiteLLMLogging(
logging_obj = get_litellm_logging_class()( # Victim for object pool
model=model, # type: ignore
messages=messages,
stream=stream,
@ -903,7 +906,7 @@ def client(original_function): # noqa: PLR0915
rules_obj = Rules()
def check_coroutine(value) -> bool:
return coroutine_checker.is_async_callable(value)
return get_coroutine_checker().is_async_callable(value)
async def async_pre_call_deployment_hook(kwargs: Dict[str, Any], call_type: str):
"""
@ -1597,7 +1600,7 @@ def client(original_function): # noqa: PLR0915
setattr(e, "timeout", timeout)
raise e
is_coroutine = coroutine_checker.is_async_callable(original_function)
is_coroutine = get_coroutine_checker().is_async_callable(original_function)
# Return the appropriate wrapper based on the original function type
if is_coroutine:

View file

@ -296,6 +296,66 @@
"output_cost_per_token": 0.0,
"output_vector_size": 1024
},
"twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"us.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"eu.twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"input_cost_per_video_per_second": 0.0007,
"input_cost_per_audio_per_second": 0.00014,
"input_cost_per_image": 0.0001,
"litellm_provider": "bedrock",
"max_input_tokens": 77,
"max_tokens": 77,
"mode": "embedding",
"output_cost_per_token": 0.0,
"output_vector_size": 1024,
"supports_embedding_image_input": true,
"supports_image_input": true
},
"twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"us.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"eu.twelvelabs.pegasus-1-2-v1:0": {
"input_cost_per_video_per_second": 0.00049,
"output_cost_per_token": 7.5e-06,
"litellm_provider": "bedrock",
"mode": "chat",
"supports_video_input": true
},
"amazon.titan-text-express-v1": {
"input_cost_per_token": 1.3e-06,
"litellm_provider": "bedrock",

View file

@ -0,0 +1,43 @@
## Tests that chat_completion endpoint has no imports inside function bodies
## This is critical for performance optimization in the hot path
import ast
from pathlib import Path
def test_chat_completion_no_imports():
"""Test that chat_completion endpoint has no imports in function bodies."""
# Path to the proxy server file
proxy_server_path = Path(__file__).parent.parent.parent / "litellm" / "proxy" / "proxy_server.py"
with open(proxy_server_path, 'r') as f:
content = f.read()
# Parse the AST
tree = ast.parse(content)
# Find the chat_completion function
chat_completion_func = None
for node in ast.walk(tree):
if (isinstance(node, ast.AsyncFunctionDef) and node.name == "chat_completion"):
chat_completion_func = node
break
assert chat_completion_func is not None, "chat_completion function not found"
# Check for imports inside the function body
import_violations = []
for node in ast.walk(chat_completion_func):
if isinstance(node, (ast.Import, ast.ImportFrom)):
# Get line number
line_num = node.lineno
import_violations.append(line_num)
# Assert no import violations found
if import_violations:
print(f"Found {len(import_violations)} import violations in chat_completion endpoint:")
for line_num in import_violations:
print(f" - Line {line_num}: Import statement found")
print("\nchat_completion endpoint should not contain imports for optimal performance.")
raise Exception("Import violations found in chat_completion endpoint")

View file

@ -76,3 +76,90 @@ def test_bedrock_embedding_models(model, input_type, embed_response):
except Exception as e:
pytest.fail(f"Error occurred: {e}")
def test_e2e_bedrock_embedding():
"""
Test text embedding with TwelveLabs Marengo.
Validates that the transformation properly extracts embedding data from TwelveLabs response format.
"""
print("Testing text embedding...")
litellm._turn_on_debug()
response = litellm.embedding(
model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0",
input=["Hello world from LiteLLM with TwelveLabs Marengo!"],
aws_region_name="us-east-1"
)
# Validate response structure
assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type"
assert hasattr(response, 'data'), "Response should have 'data' attribute"
assert len(response.data) > 0, "Response data should not be empty"
# Validate first embedding
embedding_obj = response.data[0]
assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute"
assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats"
assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty"
assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric"
# Validate embedding properties
assert embedding_obj.index == 0, "First embedding should have index 0"
assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'"
# Validate usage information
assert hasattr(response, 'usage'), "Response should have usage information"
assert response.usage is not None, "Usage should not be None"
assert response.usage.total_tokens >= 0, "Total tokens should be non-negative"
print(f"Text embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}")
def test_e2e_bedrock_embedding_image_twelvelabs_marengo():
"""
Test image embedding with TwelveLabs Marengo.
Validates that the transformation properly extracts embedding data from TwelveLabs response format for images.
"""
print("Testing image embedding...")
litellm._turn_on_debug()
# Load duck.png and convert to base64
duck_img_path = os.path.join(os.path.dirname(__file__), "duck.png")
with open(duck_img_path, "rb") as img_file:
duck_img_data = base64.b64encode(img_file.read()).decode('utf-8')
duck_img_base64 = f"data:image/png;base64,{duck_img_data}"
response = litellm.embedding(
model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0",
input=[duck_img_base64],
aws_region_name="us-east-1"
)
# Validate response structure
assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type"
assert hasattr(response, 'data'), "Response should have 'data' attribute"
assert len(response.data) > 0, "Response data should not be empty"
# Validate first embedding
embedding_obj = response.data[0]
assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute"
assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats"
assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty"
assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric"
# Validate embedding properties
assert embedding_obj.index == 0, "First embedding should have index 0"
assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'"
# Validate usage information
assert hasattr(response, 'usage'), "Response should have usage information"
assert response.usage is not None, "Usage should not be None"
assert response.usage.total_tokens >= 0, "Total tokens should be non-negative"
# TwelveLabs Marengo should return 1024-dimensional embeddings
expected_dimension = 1024
assert len(embedding_obj.embedding) == expected_dimension, f"TwelveLabs Marengo should return {expected_dimension}-dimensional embeddings, got {len(embedding_obj.embedding)}"
print(f"Image embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}")

View file

@ -254,10 +254,17 @@ async def test_cohere_request_body_with_allowed_params():
}
}]
client = AsyncHTTPHandler()
# Create a mock response
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"text": "I am Command, a language model developed by Cohere.",
"generation_id": "mock-generation-id",
"finish_reason": "COMPLETE"
}
# Mock the post method
with patch.object(client, "post", new=AsyncMock()) as mock_post:
# Mock the AsyncHTTPHandler.post method at the module level
with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", return_value=mock_response) as mock_post:
try:
await litellm.acompletion(
model="cohere/command",
@ -265,8 +272,7 @@ async def test_cohere_request_body_with_allowed_params():
allowed_openai_params=["tools", "response_format", "reasoning_effort"],
response_format=test_response_format,
reasoning_effort=test_reasoning_effort,
tools=test_tools,
client=client
tools=test_tools
)
except Exception:
pass # We only care about the request body validation

View file

@ -3026,10 +3026,13 @@ def test_custom_api_base(api_base):
stream=stream,
auth_header=None,
url="my-fake-endpoint",
model="gemini-1.5-pro", # Required for Gemini custom API base URLs
)
if api_base:
assert url == api_base + ":"
# For Gemini with custom API base, URL should be constructed as api_base/models/model:endpoint
expected_url = f"{api_base}/models/gemini-1.5-pro:"
assert url == expected_url
else:
assert url == test_endpoint

View file

@ -252,8 +252,8 @@ async def create_test_team(
async def create_test_user(
session: aiohttp.ClientSession, user_data: Dict[str, Any]
) -> str:
"""Create a new user and return the user_id"""
) -> Dict[str, Any]:
"""Create a new user and return the user info"""
url = "http://0.0.0.0:4000/user/new"
headers = {
"Authorization": "Bearer sk-1234",
@ -576,10 +576,10 @@ async def test_user_email_in_all_required_metrics():
Test that user_email label is present in all the metrics that were requested to have it:
- litellm_proxy_total_requests_metric_total
- litellm_proxy_failed_requests_metric_total
- litellm_input_tokens_total
- litellm_output_tokens_total
- litellm_input_tokens_metric_total
- litellm_output_tokens_metric_total
- litellm_requests_metric_total
- litellm_spend_metric_total
- litellm_spend_metric
"""
async with aiohttp.ClientSession() as session:
# Create a user with user_email
@ -608,15 +608,15 @@ async def test_user_email_in_all_required_metrics():
# Check that user_email appears in all the required metrics
required_metrics_with_user_email = [
"litellm_proxy_total_requests_metric_total",
"litellm_input_tokens_total",
"litellm_output_tokens_total",
"litellm_input_tokens_metric_total",
"litellm_output_tokens_metric_total",
"litellm_requests_metric_total",
"litellm_spend_metric_total"
"litellm_spend_metric"
]
import re
for metric_name in required_metrics_with_user_email:
# Check that the metric exists and contains user_email label
import re
# Look for the metric with user_email in its labels
pattern = rf'{metric_name}{{[^}}]*user_email="{re.escape(user_email)}"[^}}]*}}'
matches = re.findall(pattern, metrics_text)

View file

@ -3428,6 +3428,16 @@ async def test_list_keys(prisma_client):
),
page=1,
size=10,
user_id=None,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=None,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
print("response=", response)
assert "keys" in response
@ -3442,6 +3452,16 @@ async def test_list_keys(prisma_client):
UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value),
page=1,
size=2,
user_id=None,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=None,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
print("pagination response=", response)
assert len(response["keys"]) == 2
@ -3470,9 +3490,18 @@ async def test_list_keys(prisma_client):
response = await list_keys(
request,
UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value),
user_id=user_id,
page=1,
size=10,
user_id=user_id,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=None,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
print("filtered user_id response=", response)
assert len(response["keys"]) == 1
@ -3482,9 +3511,18 @@ async def test_list_keys(prisma_client):
response = await list_keys(
request,
UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value),
key_alias=key_alias,
page=1,
size=10,
user_id=None,
team_id=None,
organization_id=None,
key_hash=None,
key_alias=key_alias,
return_full_object=False,
include_team_keys=False,
include_created_by_keys=False,
sort_by=None,
sort_order="desc",
)
assert len(response["keys"]) == 1
assert _key in response["keys"]

View file

@ -203,6 +203,9 @@ class TestDataDogLLMObsLogger:
assert metadata["cache_hit"] is True
assert metadata["cache_key"] == "test-cache-key-789"
# Test 4: Verify is_streamed_request is in metadata
assert metadata["is_streamed_request"] is True
def test_cache_metadata_fields(self, mock_env_vars, mock_response_obj):
"""Test that cache-related metadata fields are correctly tracked"""
with patch(

View file

@ -1597,7 +1597,7 @@ async def test_no_cache_control_no_cache_point():
# ============================================================================
def test_guarded_text_wraps_in_guardrail_converse_content():
"""Test that guarded_text content type gets wrapped in guardrailConverseContent blocks."""
"""Test that guarded_text content type gets wrapped in guardContent blocks."""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_converse_messages_pt
messages = [
@ -1631,9 +1631,9 @@ def test_guarded_text_wraps_in_guardrail_converse_content():
assert "text" in content[2]
assert content[2]["text"] == "More regular text"
# Second should be guardrailConverseContent
assert "guardrailConverseContent" in content[1]
assert content[1]["guardrailConverseContent"]["text"] == "This should be guarded"
# Second should be guardContent
assert "guardContent" in content[1]
assert content[1]["guardContent"]["text"]["text"] == "This should be guarded"
def test_guarded_text_with_system_messages():
@ -1685,9 +1685,9 @@ def test_guarded_text_with_system_messages():
assert "text" in content[0]
assert content[0]["text"] == "What is the main topic of this legal document?"
# Second should be guardrailConverseContent
assert "guardrailConverseContent" in content[1]
assert content[1]["guardrailConverseContent"]["text"] == "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question."
# Second should be guardContent
assert "guardContent" in content[1]
assert content[1]["guardContent"]["text"]["text"] == "This is a set of very long instructions that you will follow. Here is a legal document that you will use to answer the user's question."
def test_guarded_text_with_mixed_content_types():
@ -1726,9 +1726,9 @@ def test_guarded_text_with_mixed_content_types():
# Second should be image
assert "image" in content[1]
# Third should be guardrailConverseContent
assert "guardrailConverseContent" in content[2]
assert content[2]["guardrailConverseContent"]["text"] == "This sensitive content should be guarded"
# Third should be guardContent
assert "guardContent" in content[2]
assert content[2]["guardContent"]["text"]["text"] == "This sensitive content should be guarded"
@pytest.mark.asyncio
@ -1764,9 +1764,9 @@ async def test_async_guarded_text():
assert "text" in content[0]
assert content[0]["text"] == "Hello"
# Second should be guardrailConverseContent
assert "guardrailConverseContent" in content[1]
assert content[1]["guardrailConverseContent"]["text"] == "This should be guarded"
# Second should be guardContent
assert "guardContent" in content[1]
assert content[1]["guardContent"]["text"]["text"] == "This should be guarded"
def test_guarded_text_with_tool_calls():
@ -1818,15 +1818,15 @@ def test_guarded_text_with_tool_calls():
assert "text" in content[0]
assert content[0]["text"] == "What's the weather?"
# Second should be guardrailConverseContent
assert "guardrailConverseContent" in content[1]
assert content[1]["guardrailConverseContent"]["text"] == "Please be careful with sensitive information"
# Second should be guardContent
assert "guardContent" in content[1]
assert content[1]["guardContent"]["text"]["text"] == "Please be careful with sensitive information"
# Other messages should not have guardrailConverseContent
# Other messages should not have guardContent
for i in range(1, 3):
content = result[i]["content"]
for block in content:
assert "guardrailConverseContent" not in block
assert "guardContent" not in block
def test_guarded_text_guardrail_config_preserved():
@ -2066,234 +2066,11 @@ def test_auto_convert_in_full_transformation():
assert "messages" in result
assert len(result["messages"]) == 1
# The message should have guardrailConverseContent
# The message should have guardContent
message = result["messages"][0]
assert "content" in message
assert len(message["content"]) == 1
assert "guardrailConverseContent" in message["content"][0]
assert message["content"][0]["guardrailConverseContent"]["text"] == "What is the main topic of this legal document?"
assert "guardContent" in message["content"][0]
assert message["content"][0]["guardContent"]["text"]["text"] == "What is the main topic of this legal document?"
def test_convert_consecutive_user_messages_to_guarded_text():
"""Test that consecutive user messages at the end are converted to guarded_text."""
config = AmazonConverseConfig()
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"text": "First user message"
}
]
},
{
"role": "assistant",
"content": "Assistant response"
},
{
"role": "user",
"content": [
{
"type": "text",
"text": "Second user message"
}
]
},
{
"role": "user",
"content": [
{
"type": "text",
"text": "Third user message"
}
]
}
]
optional_params = {
"guardrailConfig": {
"guardrailIdentifier": "gr-abc123",
"guardrailVersion": "1"
}
}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
# Verify the conversion - only the last two user messages should be converted
assert len(converted_messages) == 4
# First user message should remain unchanged
assert converted_messages[0]["role"] == "user"
assert converted_messages[0]["content"][0]["type"] == "text"
assert converted_messages[0]["content"][0]["text"] == "First user message"
# Assistant message should remain unchanged
assert converted_messages[1]["role"] == "assistant"
assert converted_messages[1]["content"] == "Assistant response"
# Second user message should be converted to guarded_text
assert converted_messages[2]["role"] == "user"
assert converted_messages[2]["content"][0]["type"] == "guarded_text"
assert converted_messages[2]["content"][0]["text"] == "Second user message"
# Third user message should be converted to guarded_text
assert converted_messages[3]["role"] == "user"
assert converted_messages[3]["content"][0]["type"] == "guarded_text"
assert converted_messages[3]["content"][0]["text"] == "Third user message"
def test_convert_all_user_messages_when_all_consecutive():
"""Test that all user messages are converted when they are all consecutive at the end."""
config = AmazonConverseConfig()
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"text": "First user message"
}
]
},
{
"role": "user",
"content": [
{
"type": "text",
"text": "Second user message"
}
]
},
{
"role": "user",
"content": [
{
"type": "text",
"text": "Third user message"
}
]
}
]
optional_params = {
"guardrailConfig": {
"guardrailIdentifier": "gr-abc123",
"guardrailVersion": "1"
}
}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
# Verify all three user messages are converted
assert len(converted_messages) == 3
for i in range(3):
assert converted_messages[i]["role"] == "user"
assert converted_messages[i]["content"][0]["type"] == "guarded_text"
assert converted_messages[0]["content"][0]["text"] == "First user message"
assert converted_messages[1]["content"][0]["text"] == "Second user message"
assert converted_messages[2]["content"][0]["text"] == "Third user message"
def test_convert_consecutive_user_messages_with_string_content():
"""Test that consecutive user messages with string content are converted to guarded_text."""
config = AmazonConverseConfig()
messages = [
{
"role": "assistant",
"content": "Assistant response"
},
{
"role": "user",
"content": "First user message"
},
{
"role": "user",
"content": "Second user message"
}
]
optional_params = {
"guardrailConfig": {
"guardrailIdentifier": "gr-abc123",
"guardrailVersion": "1"
}
}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
# Verify the conversion
assert len(converted_messages) == 3
# Assistant message should remain unchanged
assert converted_messages[0]["role"] == "assistant"
assert converted_messages[0]["content"] == "Assistant response"
# Both user messages should be converted to guarded_text
assert converted_messages[1]["role"] == "user"
assert len(converted_messages[1]["content"]) == 1
assert converted_messages[1]["content"][0]["type"] == "guarded_text"
assert converted_messages[1]["content"][0]["text"] == "First user message"
assert converted_messages[2]["role"] == "user"
assert len(converted_messages[2]["content"]) == 1
assert converted_messages[2]["content"][0]["type"] == "guarded_text"
assert converted_messages[2]["content"][0]["text"] == "Second user message"
def test_skip_consecutive_user_messages_with_existing_guarded_text():
"""Test that consecutive user messages with existing guarded_text are skipped."""
config = AmazonConverseConfig()
messages = [
{
"role": "user",
"content": [
{
"type": "guarded_text",
"text": "Already guarded"
}
]
},
{
"role": "user",
"content": [
{
"type": "text",
"text": "Should be converted"
}
]
}
]
optional_params = {
"guardrailConfig": {
"guardrailIdentifier": "gr-abc123",
"guardrailVersion": "1"
}
}
# Test the helper method directly
converted_messages = config._convert_consecutive_user_messages_to_guarded_text(messages, optional_params)
# Verify the conversion
assert len(converted_messages) == 2
# First message should remain unchanged (already has guarded_text)
assert converted_messages[0]["role"] == "user"
assert converted_messages[0]["content"][0]["type"] == "guarded_text"
assert converted_messages[0]["content"][0]["text"] == "Already guarded"
# Second message should be converted
assert converted_messages[1]["role"] == "user"
assert converted_messages[1]["content"][0]["type"] == "guarded_text"
assert converted_messages[1]["content"][0]["text"] == "Should be converted"

View file

@ -20,6 +20,13 @@ cohere_embedding_response = {
"inputTextTokenCount": 10
}
twelvelabs_embedding_response = {
"embedding": [0.1, 0.2, 0.3],
"embeddingOption": "visual-text",
"startSec": 0.0,
"endSec": 1.0
}
# Test data
test_input = "Hello world from litellm"
test_image_base64 = "data:image/png,test_image_base64_data"
@ -33,6 +40,8 @@ test_image_base64 = "data:image/png,test_image_base64_data"
("bedrock/amazon.titan-embed-image-v1", "image", titan_embedding_response),
("bedrock/cohere.embed-english-v3", "text", cohere_embedding_response),
("bedrock/cohere.embed-multilingual-v3", "text", cohere_embedding_response),
("bedrock/twelvelabs.marengo-embed-2-7-v1:0", "text", twelvelabs_embedding_response),
("bedrock/twelvelabs.marengo-embed-2-7-v1:0", "image", twelvelabs_embedding_response),
],
)
def test_bedrock_embedding_with_api_key_bearer_token(model, input_type, embed_response):