mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge branch 'litellm_contributor_prs_09_18_2025_p2' into fix/issue-14685-bedrock-titan-v2-encoding-format
This commit is contained in:
commit
63c26d7a4f
27 changed files with 876 additions and 363 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
95
docs/my-website/docs/providers/bedrock_embedding.md
Normal file
95
docs/my-website/docs/providers/bedrock_embedding.md
Normal 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)
|
||||
|
|
@ -411,6 +411,7 @@ const sidebars = {
|
|||
label: "Bedrock",
|
||||
items: [
|
||||
"providers/bedrock",
|
||||
"providers/bedrock_embedding",
|
||||
"providers/bedrock_agents",
|
||||
"providers/bedrock_batches",
|
||||
"providers/bedrock_vector_store",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
56
litellm/litellm_core_utils/cached_imports.py
Normal file
56
litellm/litellm_core_utils/cached_imports.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
140
litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py
Normal file
140
litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py
Normal 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)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
43
tests/code_coverage_tests/test_chat_completion_imports.py
Normal file
43
tests/code_coverage_tests/test_chat_completion_imports.py
Normal 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")
|
||||
|
|
@ -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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue