From 544db8d140533df4e6cc5c8e3884f310540998ac Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 3 Oct 2025 07:22:33 +0530 Subject: [PATCH] (feat)Litellm x twelvelabs bedrock[Async Invoke Support] (#14871) * Add async invoke support * Add docs and correct embedding response * fix cicd erros * fix cicd erros * fix mypy error * Add litellm param input_type * Update the docs --- .../docs/embedding/supported_embedding.md | 52 +++ .../docs/providers/bedrock_embedding.md | 176 +++++++++ litellm/__init__.py | 1 + litellm/batches/main.py | 144 +++++++- litellm/constants.py | 2 +- litellm/llms/bedrock/common_utils.py | 130 ++++--- litellm/llms/bedrock/embed/embedding.py | 225 ++++++++++-- .../twelvelabs_marengo_transformation.py | 202 +++++++++-- litellm/types/llms/bedrock.py | 33 +- litellm/types/utils.py | 25 +- litellm/utils.py | 2 + .../llm_translation/test_bedrock_embedding.py | 86 +++++ .../test_bedrock_async_invoke_embedding.py | 336 ++++++++++++++++++ .../bedrock/embed/test_bedrock_embedding.py | 177 ++++++++- 14 files changed, 1445 insertions(+), 146 deletions(-) create mode 100644 tests/test_litellm/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py diff --git a/docs/my-website/docs/embedding/supported_embedding.md b/docs/my-website/docs/embedding/supported_embedding.md index 1fd5a03e652..e63d9403665 100644 --- a/docs/my-website/docs/embedding/supported_embedding.md +++ b/docs/my-website/docs/embedding/supported_embedding.md @@ -266,7 +266,59 @@ print(response) | Titan Embeddings - G1 | `embedding(model="amazon.titan-embed-text-v1", input=input)` | | Cohere Embeddings - English | `embedding(model="cohere.embed-english-v3", input=input)` | | Cohere Embeddings - Multilingual | `embedding(model="cohere.embed-multilingual-v3", input=input)` | +| TwelveLabs Marengo (Async) | `embedding(model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", input=input, input_type="text")` | [Async Invoke Docs](../providers/bedrock_embedding#async-invoke-embedding) | +## TwelveLabs Bedrock Embedding Models + +TwelveLabs Marengo models support multimodal embeddings (text, image, video, audio) and require the `input_type` parameter to specify the input format. + +### Usage + +```python +from litellm import embedding +import os + +# Set AWS credentials +os.environ["AWS_ACCESS_KEY_ID"] = "" +os.environ["AWS_SECRET_ACCESS_KEY"] = "" +os.environ["AWS_REGION_NAME"] = "us-east-1" + +# Text embedding +response = embedding( + model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["Hello world from LiteLLM!"], + input_type="text" # Required parameter +) + +# Image embedding (base64) +response = embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["data:image/jpeg;base64,/9j/4AAQSkZJRgABAQAAAQ..."], + input_type="image", # Required parameter + output_s3_uri="s3://your-bucket/async-invoke-output/" +) + +# Video embedding (S3 URL) +response = embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["s3://your-bucket/video.mp4"], + input_type="video", # Required parameter + output_s3_uri="s3://your-bucket/async-invoke-output/" +) +``` + +### Required Parameters + +| Parameter | Description | Values | +|-----------|-------------|--------| +| `input_type` | Type of input content | `"text"`, `"image"`, `"video"`, `"audio"` | + +### Supported Models + +| Model Name | Function Call | Notes | +|------------|---------------|-------| +| TwelveLabs Marengo 2.7 (Sync) | `embedding(model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0", input=input, input_type="text")` | Text embeddings only | +| TwelveLabs Marengo 2.7 (Async) | `embedding(model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", input=input, input_type="text/image/video/audio")` | All input types, requires `output_s3_uri` | ## Cohere Embedding Models https://docs.cohere.com/reference/embed diff --git a/docs/my-website/docs/providers/bedrock_embedding.md b/docs/my-website/docs/providers/bedrock_embedding.md index 95ee8d3d228..69c5f3c86ec 100644 --- a/docs/my-website/docs/providers/bedrock_embedding.md +++ b/docs/my-website/docs/providers/bedrock_embedding.md @@ -8,6 +8,182 @@ | 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) | +## Async Invoke Support + +LiteLLM supports AWS Bedrock's async-invoke feature for embedding models that require asynchronous processing, particularly useful for large media files (video, audio) or when you need to process embeddings in the background. + +### Supported Models + +| Provider | Async Invoke Route | Use Case | +|----------|-------------------|----------| +| TwelveLabs Marengo | `bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0` | Video, audio, image, and text embeddings | + +### Required Parameters + +When using async-invoke, you must provide: + +| Parameter | Description | Required | +|-----------|-------------|----------| +| `output_s3_uri` | S3 URI where the embedding results will be stored | ✅ Yes | +| `input_type` | Type of input: `"text"`, `"image"`, `"video"`, or `"audio"` | ✅ Yes | +| `aws_region_name` | AWS region for the request | ✅ Yes | + +### Usage + +#### Basic Async Invoke + +```python +from litellm import embedding + +# Text embedding with async-invoke +response = embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["Hello world from LiteLLM async invoke!"], + aws_region_name="us-east-1", + input_type="text", + output_s3_uri="s3://your-bucket/async-invoke-output/" +) + +print(f"Job submitted! Invocation ARN: {response._hidden_params._invocation_arn}") +``` + +#### Video/Audio Embedding + +```python +# Video embedding (requires async-invoke) +response = embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["s3://your-bucket/video.mp4"], # S3 URL for video + aws_region_name="us-east-1", + input_type="video", + output_s3_uri="s3://your-bucket/async-invoke-output/" +) + +print(f"Video embedding job submitted! ARN: {response._hidden_params._invocation_arn}") +``` + +#### Image Embedding with Base64 + +```python +import base64 + +# Load and encode image +with open("image.jpg", "rb") as img_file: + img_data = base64.b64encode(img_file.read()).decode('utf-8') + img_base64 = f"data:image/jpeg;base64,{img_data}" + +response = embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=[img_base64], + aws_region_name="us-east-1", + input_type="image", + output_s3_uri="s3://your-bucket/async-invoke-output/" +) +``` + +### Retrieving Job Information + +#### Getting Job ID and Invocation ARN + +The async-invoke response includes the invocation ARN in the hidden parameters: + +```python +response = embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["Hello world"], + aws_region_name="us-east-1", + input_type="text", + output_s3_uri="s3://your-bucket/async-invoke-output/" +) + +# Access invocation ARN +invocation_arn = response._hidden_params._invocation_arn +print(f"Invocation ARN: {invocation_arn}") + +# Extract job ID from ARN (last part after the last slash) +job_id = invocation_arn.split("/")[-1] +print(f"Job ID: {job_id}") +``` + +#### Checking Job Status + +Use LiteLLM's `retrieve_batch` function to check if your job is still processing: + +```python +from litellm import retrieve_batch + +def check_async_job_status(invocation_arn, aws_region_name="us-east-1"): + """Check the status of an async invoke job using LiteLLM batch API""" + try: + response = retrieve_batch( + batch_id=invocation_arn, + custom_llm_provider="bedrock", + aws_region_name=aws_region_name + ) + return response + except Exception as e: + print(f"Error checking job status: {e}") + return None + +# Check status +status = check_async_job_status(invocation_arn, "us-east-1") +if status: + print(f"Job Status: {status.status}") + print(f"Output Location: {status.output_file_id}") +``` + +**Note:** The actual embedding results are stored in S3. The `output_file_id` from the batch status can be used to locate the results file in your S3 bucket. + +### Error Handling + +#### Common Errors + +| Error | Cause | Solution | +|-------|-------|----------| +| `ValueError: output_s3_uri cannot be empty` | Missing S3 output URI | Provide a valid S3 URI | +| `ValueError: Input type 'video' requires async_invoke route` | Using video/audio without async-invoke | Use `bedrock/async_invoke/` model prefix | +| `ValueError: input_type is required` | Missing input type parameter | Specify `input_type` parameter | + +#### Example Error Handling + +```python +try: + response = embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["Hello world"], + aws_region_name="us-east-1", + input_type="text", + output_s3_uri="s3://your-bucket/output/" # Required for async-invoke + ) + print("Job submitted successfully!") + +except ValueError as e: + if "output_s3_uri cannot be empty" in str(e): + print("Error: Please provide a valid S3 output URI") + elif "requires async_invoke route" in str(e): + print("Error: Use async_invoke model for video/audio inputs") + else: + print(f"Error: {e}") +except Exception as e: + print(f"Unexpected error: {e}") +``` + +### Best Practices + +1. **Use async-invoke for large files**: Video and audio files are better processed asynchronously +2. **Use LiteLLM batch API**: Use `retrieve_batch()` instead of direct Bedrock API calls for status checking +3. **Monitor job status**: Check job status periodically using the batch API to know when results are ready +4. **Handle errors gracefully**: Implement proper error handling for network issues and job failures +5. **Set appropriate timeouts**: Consider the processing time for large files +6. **Use S3 for large inputs**: For video/audio, use S3 URLs instead of base64 encoding + +### Limitations + +- Async-invoke is currently only supported for TwelveLabs Marengo models +- Results are stored in S3 and must be retrieved separately using the output file ID +- Job status checking requires using LiteLLM's `retrieve_batch()` function +- No built-in polling mechanism in LiteLLM (must implement your own status checking loop) + ### API keys This can be set as env variables or passed as **params to litellm.embedding()** ```python diff --git a/litellm/__init__.py b/litellm/__init__.py index d961f42efde..d1f00d3f0d0 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1161,6 +1161,7 @@ from .llms.bedrock.embed.amazon_titan_v2_transformation import ( ) from .llms.cohere.chat.transformation import CohereChatConfig from .llms.bedrock.embed.cohere_transformation import BedrockCohereEmbeddingConfig +from .llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig from .llms.openai.openai import OpenAIConfig, MistralEmbeddingConfig from .llms.openai.image_variations.transformation import OpenAIImageVariationConfig from .llms.deepinfra.chat.transformation import DeepInfraConfig diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 37b9aff4efb..48521e5fba0 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -59,18 +59,22 @@ def _resolve_timeout( ) -> float: """ Resolve timeout value from various sources and handle httpx.Timeout objects. - + Args: optional_params: GenericLiteLLMParams object containing timeout kwargs: Additional kwargs that may contain request_timeout custom_llm_provider: Provider name for httpx timeout support check default_timeout: Default timeout value to use - + Returns: Resolved timeout as float """ - timeout = optional_params.timeout or kwargs.get("request_timeout", default_timeout) or default_timeout - + timeout = ( + optional_params.timeout + or kwargs.get("request_timeout", default_timeout) + or default_timeout + ) + # Handle httpx.Timeout objects if isinstance(timeout, httpx.Timeout): if supports_httpx_timeout(custom_llm_provider) is False: @@ -81,11 +85,11 @@ def _resolve_timeout( # For providers that support httpx.Timeout, we still need to return a float # This case might need to be handled differently based on the actual use case return float(timeout.read or default_timeout) - + # Handle None case if timeout is None: return float(default_timeout) - + # Handle numeric values (int, float, string representations) return float(timeout) @@ -163,15 +167,19 @@ def create_batch( try: if model is not None: model, _, _, _ = get_llm_provider( - model=model, - custom_llm_provider=None, - ) + model=model, + custom_llm_provider=None, + ) except Exception as e: - verbose_logger.exception(f"litellm.batches.main.py::create_batch() - Error inferring custom_llm_provider - {str(e)}") - + verbose_logger.exception( + f"litellm.batches.main.py::create_batch() - Error inferring custom_llm_provider - {str(e)}" + ) + _is_async = kwargs.pop("acreate_batch", False) is True litellm_params = dict(GenericLiteLLMParams(**kwargs)) - litellm_logging_obj: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None)) + litellm_logging_obj: LiteLLMLoggingObj = cast( + LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None) + ) ### TIMEOUT LOGIC ### timeout = _resolve_timeout(optional_params, kwargs, custom_llm_provider) litellm_logging_obj.update_environment_variables( @@ -189,7 +197,6 @@ def create_batch( }, custom_llm_provider=custom_llm_provider, ) - _create_batch_request = CreateBatchRequest( completion_window=completion_window, @@ -378,6 +385,7 @@ async def aretrieve_batch( except Exception as e: raise e + def _handle_retrieve_batch_providers_without_provider_config( batch_id: str, optional_params: GenericLiteLLMParams, @@ -497,6 +505,7 @@ def _handle_retrieve_batch_providers_without_provider_config( ) return response + @client def retrieve_batch( batch_id: str, @@ -513,7 +522,9 @@ def retrieve_batch( """ try: optional_params = GenericLiteLLMParams(**kwargs) - litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None) + litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( + "litellm_logging_obj", None + ) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 litellm_params = get_litellm_params( @@ -549,7 +560,26 @@ def retrieve_batch( _is_async = kwargs.pop("aretrieve_batch", False) is True client = kwargs.get("client", None) - + + # Check if this is an async invoke ARN (different from regular batch ARN) + # Async invoke ARNs have format: arn:aws(-[^:]+)?:bedrock:[a-z0-9-]{1,20}:[0-9]{12}:async-invoke/[a-z0-9]{12} + if ( + batch_id.startswith("arn:aws") + and ":bedrock:" in batch_id + and ":async-invoke/" in batch_id + ): + # Handle async invoke status check + # Remove aws_region_name from kwargs to avoid duplicate parameter + async_kwargs = kwargs.copy() + async_kwargs.pop("aws_region_name", None) + + return _handle_async_invoke_status( + batch_id=batch_id, + aws_region_name=kwargs.get("aws_region_name", "us-east-1"), + logging_obj=litellm_logging_obj, + **async_kwargs, + ) + # Try to use provider config first (for providers like bedrock) model: Optional[str] = kwargs.get("model", None) if model is not None: @@ -559,7 +589,7 @@ def retrieve_batch( ) else: provider_config = None - + if provider_config is not None: response = base_llm_http_handler.retrieve_batch( batch_id=batch_id, @@ -568,7 +598,8 @@ def retrieve_batch( headers=extra_headers or {}, api_base=optional_params.api_base, api_key=optional_params.api_key, - logging_obj=litellm_logging_obj or LiteLLMLoggingObj( + logging_obj=litellm_logging_obj + or LiteLLMLoggingObj( model=model or "bedrock/unknown", messages=[], stream=False, @@ -586,7 +617,6 @@ def retrieve_batch( model=model, ) return response - ######################################################### # Handle providers without provider config @@ -600,7 +630,7 @@ def retrieve_batch( _is_async=_is_async, timeout=timeout, ) - + except Exception as e: raise e @@ -933,3 +963,79 @@ def cancel_batch( return response except Exception as e: raise e + + +def _handle_async_invoke_status( + batch_id: str, aws_region_name: str, logging_obj=None, **kwargs +) -> "LiteLLMBatch": + """ + Handle async invoke status check for AWS Bedrock. + + Args: + batch_id: The async invoke ARN + aws_region_name: AWS region name + **kwargs: Additional parameters + + Returns: + dict: Status information including status, output_file_id (S3 URL), etc. + """ + import asyncio + + from litellm.llms.bedrock.embed.embedding import BedrockEmbedding + + async def _async_get_status(): + # Create embedding handler instance + embedding_handler = BedrockEmbedding() + + # Get the status of the async invoke job + status_response = await embedding_handler._get_async_invoke_status( + invocation_arn=batch_id, + aws_region_name=aws_region_name, + logging_obj=logging_obj, + **kwargs, + ) + + # Transform response to a LiteLLMBatch object + from litellm.types.utils import LiteLLMBatch + + result = LiteLLMBatch( + id=status_response["invocationArn"], + object="batch", + status=status_response["status"], + created_at=status_response["submitTime"], + in_progress_at=status_response["lastModifiedTime"], + completed_at=status_response.get("endTime"), + failed_at=status_response.get("endTime") + if status_response["status"] == "failed" + else None, + request_counts={ + "total": 1, + "completed": 1 if status_response["status"] == "completed" else 0, + "failed": 1 if status_response["status"] == "failed" else 0, + }, + metadata={ + "output_file_id": status_response["outputDataConfig"][ + "s3OutputDataConfig" + ]["s3Uri"], + "failure_message": status_response.get("failureMessage"), + "model_arn": status_response["modelArn"], + }, + ) + + return result + + # Since this function is called from within an async context via run_in_executor, + # we need to create a new event loop in a thread to avoid conflicts + import concurrent.futures + + def run_in_thread(): + new_loop = asyncio.new_event_loop() + asyncio.set_event_loop(new_loop) + try: + return new_loop.run_until_complete(_async_get_status()) + finally: + new_loop.close() + + with concurrent.futures.ThreadPoolExecutor() as executor: + future = executor.submit(run_in_thread) + return future.result() diff --git a/litellm/constants.py b/litellm/constants.py index 3ff9a4b6fb0..318b23c72ce 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -374,7 +374,7 @@ OPENAI_TRANSCRIPTION_PARAMS = [ "timestamp_granularities", ] -OPENAI_EMBEDDING_PARAMS = ["dimensions", "encoding_format", "user"] +OPENAI_EMBEDDING_PARAMS = ["dimensions", "encoding_format", "user", "input_type"] DEFAULT_EMBEDDING_PARAM_VALUES = { **{k: None for k in OPENAI_EMBEDDING_PARAMS}, diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 2b111cde600..89b4e1b0866 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -445,22 +445,25 @@ class BedrockModelInfo(BaseLLMModelInfo): @staticmethod def get_bedrock_route( model: str, - ) -> Literal["converse", "invoke", "converse_like", "agent"]: + ) -> Literal["converse", "invoke", "converse_like", "agent", "async_invoke"]: """ Get the bedrock route for the given model. """ - route_mappings: Dict[str, Literal["invoke", "converse_like", "converse", "agent"]] = { + route_mappings: Dict[ + str, Literal["invoke", "converse_like", "converse", "agent", "async_invoke"] + ] = { "invoke/": "invoke", - "converse_like/": "converse_like", + "converse_like/": "converse_like", "converse/": "converse", - "agent/": "agent" + "agent/": "agent", + "async_invoke/": "async_invoke", } - + # Check explicit routes first for prefix, route_type in route_mappings.items(): if prefix in model: return route_type - + base_model = BedrockModelInfo.get_base_model(model) alt_model = BedrockModelInfo.get_non_litellm_routing_model_name(model=model) if ( @@ -469,38 +472,46 @@ class BedrockModelInfo(BaseLLMModelInfo): ): return "converse" return "invoke" - + @staticmethod def _explicit_converse_route(model: str) -> bool: """ Check if the model is an explicit converse route. """ return "converse/" in model - + @staticmethod def _explicit_invoke_route(model: str) -> bool: """ Check if the model is an explicit invoke route. """ return "invoke/" in model - + @staticmethod def _explicit_agent_route(model: str) -> bool: """ Check if the model is an explicit agent route. """ return "agent/" in model - + @staticmethod def _explicit_converse_like_route(model: str) -> bool: """ Check if the model is an explicit converse like route. """ return "converse_like/" in model - @staticmethod - def get_bedrock_provider_config_for_messages_api(model: str) -> Optional[BaseAnthropicMessagesConfig]: + def _explicit_async_invoke_route(model: str) -> bool: + """ + Check if the model is an explicit async invoke route. + """ + return "async_invoke/" in model + + @staticmethod + def get_bedrock_provider_config_for_messages_api( + model: str, + ) -> Optional[BaseAnthropicMessagesConfig]: """ Get the bedrock provider config for the given model. @@ -513,19 +524,20 @@ class BedrockModelInfo(BaseLLMModelInfo): # Converse routes should go through litellm.completion() if BedrockModelInfo._explicit_converse_route(model): return None - + ######################################################### # This goes through litellm.AmazonAnthropicClaude3MessagesConfig() # Since bedrock Invoke supports Native Anthropic Messages API ######################################################### if "claude" in model: return litellm.AmazonAnthropicClaudeMessagesConfig() - + ######################################################### # These routes will go through litellm.completion() ######################################################### return None + class BedrockEventStreamDecoderBase: """ Base class for event stream decoding for Bedrock @@ -595,20 +607,20 @@ def get_anthropic_beta_from_headers(headers: dict) -> List[str]: """ Extract anthropic-beta header values and convert them to a list. Supports comma-separated values from user headers. - + Used by both converse and invoke transformations for consistent handling of anthropic-beta headers that should be passed to AWS Bedrock. - + Args: headers (dict): Request headers dictionary - + Returns: List[str]: List of anthropic beta feature strings, empty list if no header """ anthropic_beta_header = headers.get("anthropic-beta") if not anthropic_beta_header: return [] - + # Split comma-separated values and strip whitespace return [beta.strip() for beta in anthropic_beta_header.split(",")] @@ -618,19 +630,20 @@ class CommonBatchFilesUtils: Common utilities for Bedrock batch and file operations. Provides shared functionality to reduce code duplication between batches and files. """ - + def __init__(self): # Import here to avoid circular imports from .base_aws_llm import BaseAWSLLM + self._base_aws = BaseAWSLLM() def get_bedrock_model_id_from_litellm_model(self, model: str) -> str: """ Extract the actual Bedrock model ID from LiteLLM model name. - + Args: model: LiteLLM model name (e.g., "bedrock/anthropic.claude-3-sonnet-20240229-v1:0") - + Returns: Bedrock model ID (e.g., "anthropic.claude-3-sonnet-20240229-v1:0") """ @@ -641,41 +654,45 @@ class CommonBatchFilesUtils: def parse_s3_uri(self, s3_uri: str) -> tuple: """ Parse S3 URI into bucket and key components. - + Args: s3_uri: S3 URI (e.g., "s3://bucket/key/path") - + Returns: Tuple of (bucket, key) - + Raises: ValueError: If URI format is invalid """ if not s3_uri.startswith("s3://"): raise ValueError(f"Invalid S3 URI format: {s3_uri}") - + s3_parts = s3_uri[5:].split("/", 1) # Remove "s3://" and split on first "/" if len(s3_parts) != 2: raise ValueError(f"Invalid S3 URI format: {s3_uri}") - + return s3_parts[0], s3_parts[1] # bucket, key - def extract_model_from_s3_file_path(self, s3_uri: str, optional_params: dict) -> str: + def extract_model_from_s3_file_path( + self, s3_uri: str, optional_params: dict + ) -> str: """ Extract model ID from S3 file path. - + The Bedrock file transformation creates S3 objects with the model name embedded: Format: s3://bucket/litellm-bedrock-files-{model}-{uuid}.jsonl """ # Check if model is provided in optional_params first if "model" in optional_params and optional_params["model"]: - return self.get_bedrock_model_id_from_litellm_model(optional_params["model"]) - + return self.get_bedrock_model_id_from_litellm_model( + optional_params["model"] + ) + # Extract model from S3 URI path # Expected format: s3://bucket/litellm-bedrock-files-{model}-{uuid}.jsonl try: bucket, object_key = self.parse_s3_uri(s3_uri) - + # Extract model from object key if it follows our naming pattern if object_key.startswith("litellm-bedrock-files-"): # Remove prefix and suffix to get model part @@ -690,7 +707,7 @@ class CommonBatchFilesUtils: return model_name except Exception: pass - + # Fallback to default model return "anthropic.claude-3-5-sonnet-20240620-v1:0" @@ -704,14 +721,14 @@ class CommonBatchFilesUtils: ) -> tuple: """ Sign AWS request using Signature Version 4. - + Args: service_name: AWS service name ("bedrock" or "s3") data: Request data (string or dict) endpoint_url: Full endpoint URL optional_params: Optional parameters containing AWS credentials method: HTTP method (default: POST) - + Returns: Tuple of (signed_headers, signed_data) """ @@ -736,7 +753,7 @@ class CommonBatchFilesUtils: aws_web_identity_token=optional_params.get("aws_web_identity_token"), aws_sts_endpoint=optional_params.get("aws_sts_endpoint"), ) - + # Prepare the request data method_upper = method.upper() if method_upper == "GET": @@ -746,12 +763,13 @@ class CommonBatchFilesUtils: else: if isinstance(data, dict): import json + request_data = json.dumps(data) else: request_data = data # Prepare headers for non-GET requests headers = {"Content-Type": "application/json"} - + # Create AWS request and sign it sigv4 = SigV4Auth(credentials, service_name, aws_region_name) request = AWSRequest( @@ -759,45 +777,51 @@ class CommonBatchFilesUtils: ) sigv4.add_auth(request) prepped = request.prepare() - - return dict(prepped.headers), request_data.encode('utf-8') if isinstance(request_data, str) else request_data + + return ( + dict(prepped.headers), + request_data.encode("utf-8") + if isinstance(request_data, str) + else request_data, + ) def generate_unique_job_name(self, model: str, prefix: str = "litellm") -> str: """ Generate a unique job name for AWS services. AWS services often have length limits, so this creates a concise name. - + Args: model: Model name to include in the job name prefix: Prefix for the job name - + Returns: Unique job name (≤ 63 characters for Bedrock compatibility) """ from litellm._uuid import uuid + unique_id = str(uuid.uuid4())[:8] # Format: {prefix}-batch-{model}-{uuid} # Example: litellm-batch-claude-266c398e job_name = f"{prefix}-batch-{unique_id}" - + return job_name def get_s3_bucket_and_key_from_config( - self, - litellm_params: dict, + self, + litellm_params: dict, optional_params: dict, bucket_env_var: str = "AWS_S3_BUCKET_NAME", - key_prefix: str = "litellm" + key_prefix: str = "litellm", ) -> tuple: """ Get S3 bucket and generate a unique key from configuration. - + Args: litellm_params: LiteLLM parameters optional_params: Optional parameters bucket_env_var: Environment variable name for bucket key_prefix: Prefix for the S3 key - + Returns: Tuple of (bucket_name, object_key) """ @@ -806,18 +830,20 @@ class CommonBatchFilesUtils: # Get bucket name bucket_name = ( - litellm_params.get("s3_bucket_name") + litellm_params.get("s3_bucket_name") or optional_params.get("s3_bucket_name") or os.getenv(bucket_env_var) ) if not bucket_name: - raise ValueError(f"S3 bucket name is required. Set 's3_bucket_name' parameter or {bucket_env_var} env var") - + raise ValueError( + f"S3 bucket name is required. Set 's3_bucket_name' parameter or {bucket_env_var} env var" + ) + # Generate unique object key timestamp = int(time.time()) unique_id = str(uuid.uuid4())[:8] object_key = f"{key_prefix}-{timestamp}-{unique_id}" - + return bucket_name, object_key def get_error_class( @@ -827,7 +853,5 @@ class CommonBatchFilesUtils: Get Bedrock-specific error class. """ return BedrockError( - status_code=status_code, - message=error_message, - headers=headers + status_code=status_code, message=error_message, headers=headers ) diff --git a/litellm/llms/bedrock/embed/embedding.py b/litellm/llms/bedrock/embed/embedding.py index d4dd716a1f4..3edd6d6741b 100644 --- a/litellm/llms/bedrock/embed/embedding.py +++ b/litellm/llms/bedrock/embed/embedding.py @@ -22,9 +22,8 @@ from litellm.secret_managers.main import get_secret from litellm.types.llms.bedrock import ( AmazonEmbeddingRequest, CohereEmbeddingRequest, - TwelveLabsMarengoEmbeddingRequest, ) -from litellm.types.utils import EmbeddingResponse +from litellm.types.utils import EmbeddingResponse, LlmProviders from ..base_aws_llm import BaseAWSLLM from ..common_utils import BedrockError @@ -77,7 +76,7 @@ class BedrockEmbedding(BaseAWSLLM): if aws_region_name is None: aws_region_name = "us-west-2" - credentials: Credentials = self.get_credentials( + credentials: Credentials = self.get_credentials( # type: ignore aws_access_key_id=aws_access_key_id, aws_secret_access_key=aws_secret_access_key, aws_session_token=aws_session_token, @@ -151,35 +150,80 @@ 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 + self, + response_list: List[dict], + model: str, + provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL, + is_async_invoke: Optional[bool] = False, ) -> 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( + + # Handle async invoke responses (single response with invocationArn) + if ( + is_async_invoke + and len(response_list) == 1 + and "invocationArn" in response_list[0] + ): + if provider == "twelvelabs": + returned_response = ( + TwelveLabsMarengoEmbeddingConfig()._transform_async_invoke_response( + response=response_list[0], model=model + ) + ) + else: + # For other providers, create a generic async response + invocation_arn = response_list[0].get("invocationArn", "") + + from litellm.types.utils import Embedding, Usage + + embedding = Embedding( + embedding=[], + index=0, + object="embedding", # Must be literal "embedding" + ) + usage = Usage(prompt_tokens=0, total_tokens=0) + + # Create hidden params with job ID + from litellm.types.llms.base import HiddenParams + + hidden_params = HiddenParams() + setattr(hidden_params, "_invocation_arn", invocation_arn) + + returned_response = EmbeddingResponse( + data=[embedding], + model=model, + usage=usage, + hidden_params=hidden_params, + ) + else: + # Handle regular invoke responses + 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-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 - ) - - - ########################################################## + 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: @@ -203,6 +247,7 @@ class BedrockEmbedding(BaseAWSLLM): logging_obj: Any, provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL, api_key: Optional[str] = None, + is_async_invoke: Optional[bool] = False, ): responses: List[dict] = [] for data in batch_data: @@ -210,7 +255,7 @@ class BedrockEmbedding(BaseAWSLLM): if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} - prepped = self.get_request_headers( + prepped = self.get_request_headers( # type: ignore # type: ignore credentials=credentials, aws_region_name=aws_region_name, extra_headers=extra_headers, @@ -249,7 +294,10 @@ class BedrockEmbedding(BaseAWSLLM): responses.append(response) return self._transform_response( - response_list=responses, model=model, provider=provider + response_list=responses, + model=model, + provider=provider, + is_async_invoke=is_async_invoke, ) async def _async_single_func_embeddings( @@ -265,6 +313,7 @@ class BedrockEmbedding(BaseAWSLLM): logging_obj: Any, provider: BEDROCK_EMBEDDING_PROVIDERS_LITERAL, api_key: Optional[str] = None, + is_async_invoke: Optional[bool] = False, ): responses: List[dict] = [] for data in batch_data: @@ -272,7 +321,7 @@ class BedrockEmbedding(BaseAWSLLM): if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} - prepped = self.get_request_headers( + prepped = self.get_request_headers( # type: ignore # type: ignore credentials=credentials, aws_region_name=aws_region_name, extra_headers=extra_headers, @@ -311,7 +360,10 @@ class BedrockEmbedding(BaseAWSLLM): responses.append(response) ## TRANSFORM RESPONSE ## return self._transform_response( - response_list=responses, model=model, provider=provider + response_list=responses, + model=model, + provider=provider, + is_async_invoke=is_async_invoke, ) def embeddings( @@ -343,7 +395,10 @@ class BedrockEmbedding(BaseAWSLLM): model=model, model_id=unencoded_model_id, ) - + # Check async invoke needs to be used + has_async_invoke = "async_invoke/" in model + if has_async_invoke: + model = model.replace("async_invoke/", "", 1) provider = self.get_bedrock_embedding_provider(model) if provider is None: raise Exception( @@ -402,10 +457,14 @@ class BedrockEmbedding(BaseAWSLLM): elif provider == "twelvelabs": batch_data = [] for i in input: - twelvelabs_request: ( - TwelveLabsMarengoEmbeddingRequest - ) = TwelveLabsMarengoEmbeddingConfig()._transform_request( - input=i, inference_params=inference_params + twelvelabs_request = ( + TwelveLabsMarengoEmbeddingConfig()._transform_request( + input=i, + inference_params=inference_params, + async_invoke_route=has_async_invoke, + model_id=modelId, + output_s3_uri=inference_params.get("output_s3_uri"), + ) ) batch_data.append(twelvelabs_request) @@ -417,7 +476,10 @@ class BedrockEmbedding(BaseAWSLLM): ), aws_region_name=aws_region_name, ) - endpoint_url = f"{endpoint_url}/model/{modelId}/invoke" + if has_async_invoke: + endpoint_url = f"{endpoint_url}/async-invoke" + else: + endpoint_url = f"{endpoint_url}/model/{modelId}/invoke" if batch_data is not None: if aembedding: @@ -437,6 +499,7 @@ class BedrockEmbedding(BaseAWSLLM): logging_obj=logging_obj, api_key=api_key, provider=provider, + is_async_invoke=has_async_invoke, ) returned_response = self._single_func_embeddings( client=( @@ -454,6 +517,7 @@ class BedrockEmbedding(BaseAWSLLM): logging_obj=logging_obj, api_key=api_key, provider=provider, + is_async_invoke=has_async_invoke, ) if returned_response is None: raise Exception("Unable to map Bedrock request to provider") @@ -465,7 +529,7 @@ class BedrockEmbedding(BaseAWSLLM): if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} - prepped = self.get_request_headers( + prepped = self.get_request_headers( # type: ignore credentials=credentials, aws_region_name=aws_region_name, extra_headers=extra_headers, @@ -491,3 +555,94 @@ class BedrockEmbedding(BaseAWSLLM): client=client, headers=prepped.headers, # type: ignore ) + + async def _get_async_invoke_status( + self, invocation_arn: str, aws_region_name: str, logging_obj=None, **kwargs + ) -> dict: + """ + Get the status of an async invoke job using the GetAsyncInvoke operation. + + Args: + invocation_arn: The invocation ARN from the async invoke response + aws_region_name: AWS region name + **kwargs: Additional parameters (credentials, etc.) + + Returns: + dict: Status response from AWS Bedrock + """ + + # Get AWS credentials using the same method as other Bedrock methods + credentials, _ = self._load_credentials(kwargs) + + # Get the runtime endpoint + endpoint_url, _ = self.get_runtime_endpoint( + api_base=None, + aws_bedrock_runtime_endpoint=kwargs.get("aws_bedrock_runtime_endpoint"), + aws_region_name=aws_region_name, + ) + + # Construct the status check URL + status_url = f"{endpoint_url}/async-invoke/{invocation_arn}" + + # Prepare headers + headers = {"Content-Type": "application/json"} + + # Get AWS signed headers + prepped = self.get_request_headers( # type: ignore + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=None, + endpoint_url=status_url, + data="", # GET request, no body + headers=headers, + api_key=None, + ) + + # LOGGING + if logging_obj is not None: + # Create custom curl command for GET request + masked_headers = logging_obj._get_masked_headers(prepped.headers) + formatted_headers = " ".join( + [f"-H '{k}: {v}'" for k, v in masked_headers.items()] + ) + custom_curl = "\n\nGET Request Sent from LiteLLM:\n" + custom_curl += "curl -X GET \\\n" + custom_curl += f"{prepped.url} \\\n" + custom_curl += f"{formatted_headers}\n" + + logging_obj.pre_call( + input=invocation_arn, + api_key="", + additional_args={ + "complete_input_dict": {"invocation_arn": invocation_arn}, + "api_base": prepped.url, + "headers": prepped.headers, + "request_str": custom_curl, # Override with custom GET curl command + }, + ) + + # Make the GET request + client = get_async_httpx_client(llm_provider=LlmProviders.BEDROCK) + response = await client.get( + url=prepped.url, + headers=prepped.headers, + ) + + # LOGGING + if logging_obj is not None: + logging_obj.post_call( + input=invocation_arn, + api_key="", + original_response=response, + additional_args={ + "complete_input_dict": {"invocation_arn": invocation_arn} + }, + ) + + # Parse response + if response.status_code == 200: + return response.json() + else: + raise Exception( + f"Failed to get async invoke status: {response.status_code} - {response.text}" + ) diff --git a/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py b/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py index fdad8a65043..0d25440cd72 100644 --- a/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py +++ b/litellm/llms/bedrock/embed/twelvelabs_marengo_transformation.py @@ -1,33 +1,46 @@ """ -Transformation logic from OpenAI /v1/embeddings format to Bedrock TwelveLabs Marengo /invoke format. +Transformation logic from OpenAI /v1/embeddings format to Bedrock TwelveLabs Marengo /invoke and /async-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 typing import List, Optional, Union from litellm.types.llms.bedrock import ( + TwelveLabsAsyncInvokeRequest, TwelveLabsMarengoEmbeddingRequest, + TwelveLabsOutputDataConfig, + TwelveLabsS3Location, + TwelveLabsS3OutputDataConfig, ) 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. + Supports text, image, video, and audio inputs. + - InvokeModel: text and image inputs + - StartAsyncInvoke: video, audio, image, and text inputs """ def __init__(self) -> None: pass def get_supported_openai_params(self) -> List[str]: - return ["encoding_format", "textTruncate", "embeddingOption"] + return [ + "encoding_format", + "textTruncate", + "embeddingOption", + "startSec", + "lengthSec", + "useFixedLengthSec", + "minClipSec", + "input_type", + ] def map_openai_params( self, non_default_params: dict, optional_params: dict @@ -41,45 +54,140 @@ class TwelveLabsMarengoEmbeddingConfig: optional_params["textTruncate"] = v elif k == "embeddingOption": optional_params["embeddingOption"] = v + elif k == "input_type": + # Map input_type to inputType for Bedrock + optional_params["inputType"] = v + elif k in ["startSec", "lengthSec", "useFixedLengthSec", "minClipSec"]: + optional_params[k] = v return optional_params + def _extract_bucket_owner_from_params(self, inference_params: dict) -> str: + """ + Extract bucket owner from inference parameters. + """ + return inference_params.get("bucketOwner", "") + + def _is_s3_url(self, input: str) -> bool: + """Check if input is an S3 URL.""" + return input.startswith("s3://") + def _transform_request( - self, input: str, inference_params: dict - ) -> TwelveLabsMarengoEmbeddingRequest: + self, + input: str, + inference_params: dict, + async_invoke_route: bool = False, + model_id: Optional[str] = None, + output_s3_uri: Optional[str] = None, + ) -> Union[TwelveLabsMarengoEmbeddingRequest, TwelveLabsAsyncInvokeRequest]: """ - 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) + Transform OpenAI-style input to TwelveLabs Marengo format/async-invoke format. - if is_encoded: - # Image input - b64_str = get_base64_str(input) - transformed_request = TwelveLabsMarengoEmbeddingRequest( - inputType="image", mediaSource={"base64String": b64_str} - ) + Supports: + - Text inputs (for both invoke and async-invoke) + - Image inputs (for both invoke and async-invoke) + - Video inputs (async-invoke only) + - Audio inputs (async-invoke only) + - S3 URLs for all media types (async-invoke only) + """ + if inference_params.get("inputType"): + input_type = inference_params["inputType"] else: - # Text input - transformed_request = TwelveLabsMarengoEmbeddingRequest( - inputType="text", inputText=input + raise ValueError("input_type is required") + + # Validate that async-invoke is used for video/audio + if input_type in ["video", "audio"] and not async_invoke_route: + raise ValueError( + f"Input type '{input_type}' requires async_invoke route. " + f"Use model format: 'bedrock/async_invoke/model_id'" ) + transformed_request: TwelveLabsMarengoEmbeddingRequest = { + "inputType": input_type + } + + if input_type == "text": + transformed_request["inputText"] = input # Set default textTruncate if not specified if "textTruncate" not in inference_params: transformed_request["textTruncate"] = "end" + elif input_type in ["image", "video", "audio"]: + if self._is_s3_url(input): + # S3 URL input + s3_location: TwelveLabsS3Location = {"uri": input} + bucket_owner = self._extract_bucket_owner_from_params(inference_params) + if bucket_owner: + s3_location["bucketOwner"] = bucket_owner + + transformed_request["mediaSource"] = {"s3Location": s3_location} + else: + # Base64 encoded input + if input.startswith("data:"): + # Extract base64 data from data URL + b64_str = input.split(",", 1)[1] if "," in input else input + else: + # Direct base64 string + from litellm.utils import get_base64_str + b64_str = get_base64_str(input) + + transformed_request["mediaSource"] = {"base64String": b64_str} + # Apply any additional inference parameters for k, v in inference_params.items(): if k not in [ "inputType", "inputText", "mediaSource", + "bucketOwner", # Don't include bucketOwner in the request ]: # Don't override core fields transformed_request[k] = v # type: ignore + # If async invoke route, wrap in the async invoke format + if async_invoke_route and model_id: + return self._wrap_async_invoke_request( + model_input=transformed_request, + model_id=model_id, + output_s3_uri=output_s3_uri, + ) + return transformed_request + def _wrap_async_invoke_request( + self, + model_input: TwelveLabsMarengoEmbeddingRequest, + model_id: str, + output_s3_uri: Optional[str] = None, + ) -> TwelveLabsAsyncInvokeRequest: + """ + Wrap the transformed request in the correct AWS Bedrock async invoke format. + + Args: + model_input: The transformed TwelveLabs Marengo embedding request + model_id: The model identifier (without async_invoke prefix) + output_s3_uri: Optional S3 URI for output data config + + Returns: + TwelveLabsAsyncInvokeRequest: The wrapped async invoke request + """ + import urllib.parse + + # Clean the model ID + unquoted_model_id = urllib.parse.unquote(model_id) + if unquoted_model_id.startswith("async_invoke/"): + unquoted_model_id = unquoted_model_id.replace("async_invoke/", "") + + # Validate that the S3 URI is not empty + if not output_s3_uri or output_s3_uri.strip() == "": + raise ValueError("output_s3_uri cannot be empty for async invoke requests") + + return TwelveLabsAsyncInvokeRequest( + modelId=unquoted_model_id, + modelInput=model_input, + outputDataConfig=TwelveLabsOutputDataConfig( + s3OutputDataConfig=TwelveLabsS3OutputDataConfig(s3Uri=output_s3_uri) + ), + ) + def _transform_response( self, response_list: List[dict], model: str ) -> EmbeddingResponse: @@ -138,3 +246,53 @@ class TwelveLabsMarengoEmbeddingConfig: usage = Usage(prompt_tokens=total_tokens, total_tokens=total_tokens) return EmbeddingResponse(data=embeddings, model=model, usage=usage) + + def _transform_async_invoke_response( + self, response: dict, model: str + ) -> EmbeddingResponse: + """ + Transform async invoke response (invocation ARN) to OpenAI format. + + AWS async invoke returns: + { + "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123" + } + + We transform this to a job-like embedding response: + { + "object": "list", + "data": [ + { + "object": "embedding_job_id:1234567890", + "embedding": [], + "index": 0 + } + ], + "model": "model", + "usage": {} + } + """ + invocation_arn = response.get("invocationArn", "") + + # Create a placeholder embedding object for the job + embedding = Embedding( + embedding=[], # Empty embedding for async jobs + index=0, + object="embedding", + ) + + # Create usage object (empty for async jobs) + usage = Usage(prompt_tokens=0, total_tokens=0) + + # Create hidden params with job ID + from litellm.types.llms.base import HiddenParams + + hidden_params = HiddenParams() + setattr(hidden_params, "_invocation_arn", invocation_arn) + + return EmbeddingResponse( + data=[embedding], + model=model, + usage=usage, + hidden_params=hidden_params, + ) diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index cebcd0522a1..df551c5bded 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -377,9 +377,14 @@ TWELVELABS_EMBEDDING_INPUT_TYPES = Literal["text", "image", "video", "audio"] TWELVELABS_EMBEDDING_OPTIONS = Literal["visual-text", "visual-image", "audio"] +class TwelveLabsS3Location(TypedDict, total=False): + uri: str + bucketOwner: str + + class TwelveLabsMediaSource(TypedDict, total=False): base64String: str - s3Location: dict # {"uri": str, "bucketOwner": str} + s3Location: TwelveLabsS3Location class TwelveLabsMarengoEmbeddingRequest(TypedDict, total=False): @@ -401,6 +406,32 @@ class TwelveLabsMarengoEmbeddingResponse(TypedDict): endSec: float +class TwelveLabsS3OutputDataConfig(TypedDict): + s3Uri: str + + +class TwelveLabsOutputDataConfig(TypedDict): + s3OutputDataConfig: TwelveLabsS3OutputDataConfig + + +class TwelveLabsAsyncInvokeRequest(TypedDict): + modelId: str + modelInput: TwelveLabsMarengoEmbeddingRequest + outputDataConfig: TwelveLabsOutputDataConfig + + +class TwelveLabsAsyncInvokeStatusResponse(TypedDict): + invocationArn: str + modelArn: str + status: str # "InProgress" | "Completed" | "Failed" + submitTime: str + lastModifiedTime: str + endTime: Optional[str] + outputDataConfig: TwelveLabsOutputDataConfig + clientRequestToken: Optional[str] + failureMessage: Optional[str] + + AmazonEmbeddingRequest = Union[ AmazonTitanMultimodalEmbeddingRequest, AmazonTitanV2EmbeddingRequest, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index d303485b3de..4e93e167530 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -123,12 +123,18 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): max_output_tokens: Required[Optional[int]] input_cost_per_token: Required[float] input_cost_per_token_flex: Optional[float] # OpenAI flex service tier pricing - input_cost_per_token_priority: Optional[float] # OpenAI priority service tier pricing + input_cost_per_token_priority: Optional[ + float + ] # OpenAI priority service tier pricing cache_creation_input_token_cost: Optional[float] cache_creation_input_token_cost_above_1hr: Optional[float] cache_read_input_token_cost: Optional[float] - cache_read_input_token_cost_flex: Optional[float] # OpenAI flex service tier pricing - cache_read_input_token_cost_priority: Optional[float] # OpenAI priority service tier pricing + cache_read_input_token_cost_flex: Optional[ + float + ] # OpenAI flex service tier pricing + cache_read_input_token_cost_priority: Optional[ + float + ] # OpenAI priority service tier pricing input_cost_per_character: Optional[float] # only for vertex ai models input_cost_per_audio_token: Optional[float] input_cost_per_token_above_128k_tokens: Optional[float] # only for vertex ai models @@ -147,7 +153,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): output_cost_per_token_batches: Optional[float] output_cost_per_token: Required[float] output_cost_per_token_flex: Optional[float] # OpenAI flex service tier pricing - output_cost_per_token_priority: Optional[float] # OpenAI priority service tier pricing + output_cost_per_token_priority: Optional[ + float + ] # OpenAI priority service tier pricing output_cost_per_character: Optional[float] # only for vertex ai models output_cost_per_audio_token: Optional[float] output_cost_per_token_above_128k_tokens: Optional[ @@ -1417,6 +1425,9 @@ class EmbeddingResponse(OpenAIObject): model = model super().__init__(model=model, object=object, data=data, usage=usage) # type: ignore + if hidden_params: + self._hidden_params = hidden_params + def __contains__(self, key): # Define custom behavior for the 'in' operator return hasattr(self, key) @@ -2638,6 +2649,7 @@ class SpecialEnums(Enum): class ServiceTier(Enum): """Enum for service tier types used in cost calculations.""" + FLEX = "flex" PRIORITY = "priority" @@ -2684,13 +2696,14 @@ CostResponseTypes = Union[ class PriorityReservationSettings(BaseModel): """ Settings for priority-based rate limiting reservation. - + Defines what priority to assign to keys without explicit priority metadata. The priority_reservation mapping is configured separately via litellm.priority_reservation. """ + default_priority: float = Field( default=0.5, - description="Priority level to assign to API keys without explicit priority metadata. Should match a key in litellm.priority_reservation." + description="Priority level to assign to API keys without explicit priority metadata. Should match a key in litellm.priority_reservation.", ) saturation_threshold: float = Field( diff --git a/litellm/utils.py b/litellm/utils.py index 24e20c0b696..ad4d2ff11b2 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2802,6 +2802,8 @@ def get_optional_params_embeddings( # noqa: PLR0915 object = litellm.AmazonTitanV2Config() elif "cohere.embed-multilingual-v3" in model: object = litellm.BedrockCohereEmbeddingConfig() + elif "twelvelabs" in model or "marengo" in model: + object = litellm.TwelveLabsMarengoEmbeddingConfig() else: # unmapped model supported_params = [] _check_valid_arg(supported_params=supported_params) diff --git a/tests/llm_translation/test_bedrock_embedding.py b/tests/llm_translation/test_bedrock_embedding.py index 15a615b8cc4..903fd310262 100644 --- a/tests/llm_translation/test_bedrock_embedding.py +++ b/tests/llm_translation/test_bedrock_embedding.py @@ -170,6 +170,92 @@ def test_e2e_bedrock_embedding_image_twelvelabs_marengo(): print(f"Image embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}") + # Restore original region name + if original_region_name: + os.environ["AWS_REGION_NAME"] = original_region_name + + +def test_e2e_bedrock_async_invoke_embedding_twelvelabs_marengo(): + """ + Test async invoke embedding with TwelveLabs Marengo. + Validates that async invoke responses include job ID in hidden parameters. + """ + print("Testing async invoke embedding...") + original_region_name = os.environ.get("AWS_REGION_NAME") + os.environ["AWS_REGION_NAME"] = "us-east-1" + litellm._turn_on_debug() + + # Mock the HTTP call to return async invoke response + with patch("litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_sync_call") as mock_call: + mock_call.return_value = { + "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-job-123" + } + + response = litellm.embedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["Hello world from LiteLLM async invoke!"], + aws_region_name="us-east-1", + inputType="text", + output_s3_uri="s3://test-bucket/async-invoke-output/" + ) + + # Validate response structure + assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type" + assert hasattr(response, '_hidden_params'), "Response should have _hidden_params" + assert response._hidden_params is not None, "Hidden params should not be None" + + # Validate hidden params contain invocation ARN + assert hasattr(response._hidden_params, '_invocation_arn'), "Hidden params should have _invocation_arn" + assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-job-123", "Invocation ARN should be preserved" + + # Validate embedding structure + assert len(response.data) == 1, "Should have one embedding" + assert response.data[0].object == "embedding", "Embedding object should be 'embedding'" + assert response.data[0].embedding == [], "Embedding should be empty for async jobs" + + print(f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}") + + # Restore original region name + if original_region_name: + os.environ["AWS_REGION_NAME"] = original_region_name + + +@pytest.mark.asyncio +async def test_e2e_bedrock_async_invoke_embedding_async_twelvelabs_marengo(): + """ + Test async invoke embedding with async calls. + Validates that async invoke responses work with aembedding. + """ + print("Testing async invoke embedding with async calls...") + original_region_name = os.environ.get("AWS_REGION_NAME") + os.environ["AWS_REGION_NAME"] = "us-east-1" + litellm._turn_on_debug() + + # Mock the async HTTP call to return async invoke response + with patch("litellm.llms.bedrock.embed.embedding.BedrockEmbedding._make_async_call") as mock_call: + mock_call.return_value = { + "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/test-async-job-456" + } + + response = await litellm.aembedding( + model="bedrock/async_invoke/us.twelvelabs.marengo-embed-2-7-v1:0", + input=["Hello world from LiteLLM async invoke async!"], + aws_region_name="us-east-1", + inputType="text", + output_s3_uri="s3://test-bucket/async-invoke-output/" + ) + + # Validate response structure + assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type" + assert hasattr(response, '_hidden_params'), "Response should have _hidden_params" + assert response._hidden_params is not None, "Hidden params should not be None" + + # Validate hidden params contain invocation ARN + assert hasattr(response._hidden_params, '_invocation_arn'), "Hidden params should have _invocation_arn" + assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123", "Invocation ARN should be preserved" + + print(f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._invocation_arn}") + # Restore original region name if original_region_name: os.environ["AWS_REGION_NAME"] = original_region_name \ No newline at end of file diff --git a/tests/test_litellm/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py b/tests/test_litellm/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py new file mode 100644 index 00000000000..436ca6e0421 --- /dev/null +++ b/tests/test_litellm/llms/bedrock/embed/test_bedrock_async_invoke_embedding.py @@ -0,0 +1,336 @@ +import json +import os +import sys +from unittest.mock import Mock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) # Adds the parent directory to the system path +import litellm +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.types.llms.base import HiddenParams + +# Mock async invoke responses +async_invoke_response = { + "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" +} + +async_invoke_status_response = { + "status": "InProgress", + "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456", + "outputDataConfig": { + "s3OutputDataConfig": { + "s3Uri": "s3://test-bucket/async-invoke-output/" + } + } +} + +async_invoke_completed_response = { + "status": "Completed", + "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456", + "outputDataConfig": { + "s3OutputDataConfig": { + "s3Uri": "s3://test-bucket/async-invoke-output/" + } + } +} + +# Test data +test_input = "Hello world from litellm async invoke" +test_image_base64 = "data:image/png,test_image_base64_data" + + +class TestBedrockAsyncInvokeEmbedding: + """Test suite for Bedrock async-invoke embedding functionality.""" + + def test_async_invoke_response_transformation_twelvelabs(self): + """Test that async invoke responses are properly transformed with hidden params.""" + from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig + + config = TwelveLabsMarengoEmbeddingConfig() + response = config._transform_async_invoke_response(async_invoke_response, "test-model") + + # Verify response structure + assert isinstance(response, litellm.EmbeddingResponse) + assert hasattr(response, '_hidden_params') + assert response._hidden_params is not None + + # Verify hidden params contain invocation ARN + assert hasattr(response._hidden_params, '_invocation_arn') + assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + + # Verify embedding structure + assert len(response.data) == 1 + assert response.data[0].object == "embedding" + assert response.data[0].embedding == [] # Empty for async jobs + assert response.data[0].index == 0 + + def test_async_invoke_response_transformation_generic(self): + """Test that generic async invoke responses are properly transformed.""" + from litellm.llms.bedrock.embed.embedding import BedrockEmbedding + + bedrock_embedding = BedrockEmbedding() + + # Mock the transformation method + response_list = [async_invoke_response] + response = bedrock_embedding._transform_response( + response_list=response_list, + model="test-model", + provider="twelvelabs", + is_async_invoke=True + ) + + # Verify response structure + assert isinstance(response, litellm.EmbeddingResponse) + assert hasattr(response, '_hidden_params') + assert response._hidden_params is not None + + # Verify hidden params contain invocation ARN + assert hasattr(response._hidden_params, '_invocation_arn') + assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + + @pytest.mark.parametrize( + "model,input_type", + [ + ("bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0", "text"), + ("bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0", "image"), + ("bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0", "video"), + ("bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0", "audio"), + ], + ) + def test_async_invoke_twelvelabs_embedding_request_transformation(self, model, input_type): + """Test that async invoke requests are properly transformed for TwelveLabs.""" + from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig + + config = TwelveLabsMarengoEmbeddingConfig() + + # Test input based on type + if input_type == "text": + input_data = test_input + elif input_type == "image": + input_data = test_image_base64 + elif input_type in ["video", "audio"]: + input_data = "s3://test-bucket/test-file.mp4" if input_type == "video" else "s3://test-bucket/test-file.wav" + + inference_params = { + "inputType": input_type, # This will be set by the parameter mapping + "output_s3_uri": "s3://test-bucket/async-invoke-output/" + } + + transformed_request = config._transform_request( + input=input_data, + inference_params=inference_params, + async_invoke_route=True, + model_id="twelvelabs.marengo-embed-2-7-v1:0", + output_s3_uri="s3://test-bucket/async-invoke-output/" + ) + + # Verify async invoke request structure + assert "modelId" in transformed_request + assert "modelInput" in transformed_request + assert "outputDataConfig" in transformed_request + assert transformed_request["modelId"] == "twelvelabs.marengo-embed-2-7-v1:0" + assert transformed_request["outputDataConfig"]["s3OutputDataConfig"]["s3Uri"] == "s3://test-bucket/async-invoke-output/" + + def test_async_invoke_twelvelabs_embedding_with_mock(self): + """Test async invoke embedding with mocked HTTP calls.""" + litellm.set_verbose = True + client = HTTPHandler() + test_api_key = "test-bearer-token-12345" + model = "bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0" + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(async_invoke_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + response = litellm.embedding( + model=model, + input=test_input, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key=test_api_key, + input_type="text", # New input_type parameter (maps to inputType) + output_s3_uri="s3://test-bucket/async-invoke-output/" + ) + + # Verify response structure + assert isinstance(response, litellm.EmbeddingResponse) + assert hasattr(response, '_hidden_params') + assert response._hidden_params is not None + assert hasattr(response._hidden_params, '_invocation_arn') + assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + + # Verify request was made to async-invoke endpoint + request_url = mock_post.call_args.kwargs.get("url", "") + assert "/async-invoke" in request_url + + @pytest.mark.asyncio + async def test_async_invoke_twelvelabs_embedding_async_with_mock(self): + """Test async invoke embedding with async calls.""" + litellm.set_verbose = True + client = AsyncHTTPHandler() + test_api_key = "test-bearer-token-12345" + model = "bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0" + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(async_invoke_response) + mock_response.json = Mock(return_value=async_invoke_response) + mock_post.return_value = mock_response + + response = await litellm.aembedding( + model=model, + input=test_input, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key=test_api_key, + inputType="text", + output_s3_uri="s3://test-bucket/async-invoke-output/" + ) + + # Verify response structure + assert isinstance(response, litellm.EmbeddingResponse) + assert hasattr(response, '_hidden_params') + assert response._hidden_params is not None + assert hasattr(response._hidden_params, '_invocation_arn') + assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + + @pytest.mark.asyncio + async def test_async_invoke_status_checking(self): + """Test async invoke status checking functionality.""" + from litellm.llms.bedrock.embed.embedding import BedrockEmbedding + + bedrock_embedding = BedrockEmbedding() + + # Mock the async status check + with patch.object(bedrock_embedding, '_get_async_invoke_status') as mock_status: + mock_status.return_value = async_invoke_status_response + + # This would be called internally, but we can test the method directly + status_response = await bedrock_embedding._get_async_invoke_status( + invocation_arn="arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456", + aws_region_name="us-east-1" + ) + + assert status_response["status"] == "InProgress" + assert "invocationArn" in status_response + + def test_async_invoke_error_handling_missing_output_s3_uri(self): + """Test error handling when output_s3_uri is missing for async invoke.""" + from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig + + config = TwelveLabsMarengoEmbeddingConfig() + + with pytest.raises(ValueError, match="output_s3_uri cannot be empty for async invoke requests"): + config._transform_request( + input=test_input, + inference_params={"inputType": "text"}, + async_invoke_route=True, + model_id="twelvelabs.marengo-embed-2-7-v1:0", + output_s3_uri="" # Empty S3 URI should raise error + ) + + def test_async_invoke_error_handling_video_audio_without_async_route(self): + """Test error handling when video/audio input is used without async invoke route.""" + from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig + + config = TwelveLabsMarengoEmbeddingConfig() + + with pytest.raises(ValueError, match="Input type 'video' requires async_invoke route"): + config._transform_request( + input="s3://test-bucket/test-video.mp4", + inference_params={"inputType": "video"}, + async_invoke_route=False, # Should fail for video without async route + model_id="twelvelabs.marengo-embed-2-7-v1:0", + output_s3_uri="s3://test-bucket/async-invoke-output/" + ) + + def test_async_invoke_invocation_arn_preservation(self): + """Test that invocation ARN is correctly preserved in hidden params.""" + from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig + + config = TwelveLabsMarengoEmbeddingConfig() + + # Test various ARN formats + test_cases = [ + "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456", + "arn:aws:bedrock:us-west-2:987654321098:async-invoke/xyz789", + "invalid-arn", + "", + ] + + for arn in test_cases: + mock_response = {"invocationArn": arn} + response = config._transform_async_invoke_response(mock_response, "test-model") + + assert response._hidden_params._invocation_arn == arn + + def test_async_invoke_hidden_params_structure(self): + """Test that hidden params have the correct structure and can be accessed.""" + from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig + + config = TwelveLabsMarengoEmbeddingConfig() + response = config._transform_async_invoke_response(async_invoke_response, "test-model") + + # Test that hidden params can be accessed like a dictionary + assert response._hidden_params.get("_invocation_arn") == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + + # Test that hidden params can be accessed like attributes + assert response._hidden_params._invocation_arn == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + + # Test that hidden params can be accessed with bracket notation + assert response._hidden_params["_invocation_arn"] == "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + + def test_async_invoke_model_parsing(self): + """Test that async invoke models are correctly parsed.""" + from litellm.llms.bedrock.embed.embedding import BedrockEmbedding + + bedrock_embedding = BedrockEmbedding() + + # Test model parsing + test_models = [ + "bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0", + "bedrock/async_invoke/amazon.titan-embed-text-v1", + "bedrock/async_invoke/cohere.embed-english-v3", + ] + + for model in test_models: + # Check if async invoke is detected + has_async_invoke = "async_invoke/" in model + assert has_async_invoke, f"Model {model} should be detected as async invoke" + + # Check model ID extraction (remove both "bedrock/" and "async_invoke/" prefixes) + if has_async_invoke: + model_id = model.replace("bedrock/async_invoke/", "", 1) + assert model_id in [ + "twelvelabs.marengo-embed-2-7-v1:0", + "amazon.titan-embed-text-v1", + "cohere.embed-english-v3" + ] + + def test_async_invoke_endpoint_construction(self): + """Test that async invoke endpoints are correctly constructed.""" + from litellm.llms.bedrock.embed.embedding import BedrockEmbedding + + bedrock_embedding = BedrockEmbedding() + + # Mock the get_runtime_endpoint method + with patch.object(bedrock_embedding, 'get_runtime_endpoint') as mock_endpoint: + mock_endpoint.return_value = ("https://bedrock-runtime.us-east-1.amazonaws.com", None) + + # Test endpoint construction for async invoke + endpoint_url, _ = bedrock_embedding.get_runtime_endpoint( + api_base=None, + aws_bedrock_runtime_endpoint=None, + aws_region_name="us-east-1" + ) + + # For async invoke, the endpoint should be modified + async_endpoint = f"{endpoint_url}/async-invoke" + assert async_endpoint == "https://bedrock-runtime.us-east-1.amazonaws.com/async-invoke" diff --git a/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py b/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py index 081176209a8..b43ec226842 100644 --- a/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py +++ b/tests/test_litellm/llms/bedrock/embed/test_bedrock_embedding.py @@ -59,14 +59,21 @@ def test_bedrock_embedding_with_api_key_bearer_token(model, input_type, embed_re input_data = test_image_base64 if input_type == "image" else test_input - response = litellm.embedding( - model=model, - input=input_data, - client=client, - aws_region_name="us-east-1", - aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", - api_key=test_api_key - ) + # Add inputType parameter for TwelveLabs Marengo models + kwargs = { + "model": model, + "input": input_data, + "client": client, + "aws_region_name": "us-east-1", + "aws_bedrock_runtime_endpoint": "https://bedrock-runtime.us-east-1.amazonaws.com", + "api_key": test_api_key + } + + # Add input_type parameter for TwelveLabs Marengo models (maps to inputType) + if "twelvelabs.marengo-embed" in model: + kwargs["input_type"] = input_type + + response = litellm.embedding(**kwargs) assert isinstance(response, litellm.EmbeddingResponse) assert isinstance(response.data[0]['embedding'], list) @@ -241,4 +248,156 @@ def test_bedrock_titan_v2_encoding_format_base64(): # Verify that the request contains embeddingTypes: ["binary"] for base64 encoding request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) assert "embeddingTypes" in request_body - assert request_body["embeddingTypes"] == ["binary"] \ No newline at end of file + assert request_body["embeddingTypes"] == ["binary"] + + +def test_twelvelabs_input_type_parameter_mapping(): + """Test that input_type parameter is correctly mapped to inputType for TwelveLabs models""" + litellm.set_verbose = True + client = HTTPHandler() + test_api_key = "test-bearer-token-12345" + model = "bedrock/twelvelabs.marengo-embed-2-7-v1:0" + + twelvelabs_response = { + "data": [{ + "embedding": [0.1, 0.2, 0.3], + "inputTextTokenCount": 10 + }] + } + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(twelvelabs_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + # Test with input_type parameter (new LiteLLM parameter) + response = litellm.embedding( + model=model, + input=test_input, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key=test_api_key, + input_type="text" # New parameter that should map to inputType + ) + + assert isinstance(response, litellm.EmbeddingResponse) + assert isinstance(response.data[0]['embedding'], list) + assert len(response.data[0]['embedding']) == 3 + + # Verify that the request contains inputType (mapped from input_type) + request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) + assert "inputType" in request_body + assert request_body["inputType"] == "text" + assert "input_type" not in request_body # Should be mapped, not passed through + + +def test_twelvelabs_input_type_parameter_mapping_async_invoke(): + """Test that input_type parameter is correctly mapped to inputType for TwelveLabs async invoke models""" + litellm.set_verbose = True + client = HTTPHandler() + test_api_key = "test-bearer-token-12345" + model = "bedrock/async_invoke/twelvelabs.marengo-embed-2-7-v1:0" + + async_invoke_response = { + "invocationArn": "arn:aws:bedrock:us-east-1:123456789012:async-invoke/abc123def456" + } + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(async_invoke_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + # Test with input_type parameter for async invoke + response = litellm.embedding( + model=model, + input=test_input, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key=test_api_key, + output_s3_uri="s3://test-bucket/async-invoke-output/", + input_type="text" # New parameter that should map to inputType + ) + + assert isinstance(response, litellm.EmbeddingResponse) + assert hasattr(response, '_hidden_params') + assert response._hidden_params is not None + assert hasattr(response._hidden_params, '_invocation_arn') + + # Verify that the request contains inputType in modelInput (mapped from input_type) + request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}")) + assert "modelInput" in request_body + assert "inputType" in request_body["modelInput"] + assert request_body["modelInput"]["inputType"] == "text" + assert "input_type" not in request_body # Should be mapped, not passed through + + +def test_twelvelabs_missing_input_type_error(): + """Test that missing input_type parameter throws an error for TwelveLabs models but not others""" + litellm.set_verbose = True + client = HTTPHandler() + test_api_key = "test-bearer-token-12345" + + # Test TwelveLabs model - should throw error + twelvelabs_model = "bedrock/twelvelabs.marengo-embed-2-7-v1:0" + twelvelabs_response = { + "data": [{ + "embedding": [0.1, 0.2, 0.3], + "inputTextTokenCount": 10 + }] + } + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(twelvelabs_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + # Test that missing input_type throws an error for TwelveLabs + with pytest.raises(Exception) as exc_info: + litellm.embedding( + model=twelvelabs_model, + input=test_input, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key=test_api_key + # No input_type parameter - should throw an error + ) + + # Verify the error message contains the expected text + assert "input_type is required" in str(exc_info.value) + + # Test Amazon Titan model - should NOT throw error (input_type not required) + titan_model = "bedrock/amazon.titan-embed-text-v1" + titan_response = { + "embedding": [0.1, 0.2, 0.3], + "inputTextTokenCount": 10 + } + + with patch.object(client, "post") as mock_post: + mock_response = Mock() + mock_response.status_code = 200 + mock_response.text = json.dumps(titan_response) + mock_response.json = lambda: json.loads(mock_response.text) + mock_post.return_value = mock_response + + # Test that missing input_type does NOT throw an error for Amazon Titan + response = litellm.embedding( + model=titan_model, + input=test_input, + client=client, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint="https://bedrock-runtime.us-east-1.amazonaws.com", + api_key=test_api_key + # No input_type parameter - should work fine + ) + + # Should succeed without input_type + assert isinstance(response, litellm.EmbeddingResponse) \ No newline at end of file