mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
(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
This commit is contained in:
parent
9c29f35c4b
commit
544db8d140
14 changed files with 1445 additions and 146 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
@ -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"]
|
||||
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)
|
||||
Loading…
Add table
Reference in a new issue