fix(opentelemetry.py): fix issue where headers were not being split correctly + feat(bedrock): add titan image generations w/ cost tracking (#15916)

* fix(opentelemetry.py): fix issue where headers were not being split correctly

* feat(bedrock/image): Support bedrock titan image generation

Closes https://github.com/BerriAI/litellm/issues/361

* build(model_prices_and_context_window.json): track titan image gen pricing

enables cost tracking per request

* feat(amazon_titan_transformation.py): support titan image generation cost tracking

* docs: document new model

* docs: update docs to indicate cost tracking + refactor rerank into separate doc

* fix: fix mypy linting error

* fix: fix type ignore
This commit is contained in:
Krish Dholakia 2025-10-25 13:45:13 -07:00 • committed by GitHub
parent 72bbdfd3f3
commit 346e036399
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 859 additions and 439 deletions

View file

@ -1999,203 +1999,13 @@ response = embedding(
### Advanced - [Pass model/provider-specific Params](https://docs.litellm.ai/docs/completion/provider_specific_params#proxy-usage)
## Image Generation
Use this for stable diffusion, and amazon nova canvas on bedrock
See [Bedrock Image Generation](./bedrock_image_gen) for using Stable Diffusion and Amazon Nova Canvas models on Bedrock.
### Usage
## Rerank API
<Tabs>
<TabItem value="sdk" label="SDK">
```python
import os
from litellm import image_generation
os.environ["AWS_ACCESS_KEY_ID"] = ""
os.environ["AWS_SECRET_ACCESS_KEY"] = ""
os.environ["AWS_REGION_NAME"] = ""
response = image_generation(
prompt="A cute baby sea otter",
model="bedrock/stability.stable-diffusion-xl-v0",
)
print(f"response: {response}")
```
**Set optional params**
```python
import os
from litellm import image_generation
os.environ["AWS_ACCESS_KEY_ID"] = ""
os.environ["AWS_SECRET_ACCESS_KEY"] = ""
os.environ["AWS_REGION_NAME"] = ""
response = image_generation(
prompt="A cute baby sea otter",
model="bedrock/stability.stable-diffusion-xl-v0",
### OPENAI-COMPATIBLE ###
size="128x512", # width=128, height=512
### PROVIDER-SPECIFIC ### see `AmazonStabilityConfig` in bedrock.py for all params
seed=30
)
print(f"response: {response}")
```
</TabItem>
<TabItem value="proxy" label="PROXY">
1. Setup config.yaml
```yaml
model_list:
- model_name: amazon.nova-canvas-v1:0
litellm_params:
model: bedrock/amazon.nova-canvas-v1:0
aws_region_name: "us-east-1"
aws_secret_access_key: my-key # OPTIONAL - all boto3 auth params supported
aws_secret_access_id: my-id # OPTIONAL - all boto3 auth params supported
```
2. Start proxy
```bash
litellm --config /path/to/config.yaml
```
3. Test it!
```bash
curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \
-d '{
"model": "amazon.nova-canvas-v1:0",
"prompt": "A cute baby sea otter"
}'
```
</TabItem>
</Tabs>
### Using Inference Profiles with Image Generation
For AWS Bedrock Application Inference Profiles with image generation, use the `model_id` parameter to specify the inference profile ARN:
<Tabs>
<TabItem value="sdk" label="SDK">
```python
from litellm import image_generation
response = image_generation(
model="bedrock/amazon.nova-canvas-v1:0",
model_id="arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0",
prompt="A cute baby sea otter"
)
print(f"response: {response}")
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```yaml
model_list:
- model_name: nova-canvas-inference-profile
litellm_params:
model: bedrock/amazon.nova-canvas-v1:0
model_id: arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0
aws_region_name: "eu-west-1"
```
</TabItem>
</Tabs>
## Supported AWS Bedrock Image Generation Models
| Model Name | Function Call |
|----------------------|---------------------------------------------|
| Stable Diffusion 3 - v0 | `embedding(model="bedrock/stability.stability.sd3-large-v1:0", prompt=prompt)` |
| Stable Diffusion - v0 | `embedding(model="bedrock/stability.stable-diffusion-xl-v0", prompt=prompt)` |
| Stable Diffusion - v0 | `embedding(model="bedrock/stability.stable-diffusion-xl-v1", prompt=prompt)` |
## Rerank API
Use Bedrock's Rerank API in the Cohere `/rerank` format.
Supported Cohere Rerank Params
- `model` - the foundation model ARN
- `query` - the query to rerank against
- `documents` - the list of documents to rerank
- `top_n` - the number of results to return
<Tabs>
<TabItem label="SDK" value="sdk">
```python
from litellm import rerank
import os
os.environ["AWS_ACCESS_KEY_ID"] = ""
os.environ["AWS_SECRET_ACCESS_KEY"] = ""
os.environ["AWS_REGION_NAME"] = ""
response = rerank(
model="bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0", # provide the model ARN - get this here https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/bedrock/client/list_foundation_models.html
query="hello",
documents=["hello", "world"],
top_n=2,
)
print(response)
```
</TabItem>
<TabItem label="PROXY" value="proxy">
1. Setup config.yaml
```yaml
model_list:
- model_name: bedrock-rerank
litellm_params:
model: bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: os.environ/AWS_REGION_NAME
```
2. Start proxy server
```bash
litellm --config config.yaml
# RUNNING on http://0.0.0.0:4000
```
3. Test it!
```bash
curl http://0.0.0.0:4000/rerank \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"model": "bedrock-rerank",
"query": "What is the capital of the United States?",
"documents": [
"Carson City is the capital city of the American state of Nevada.",
"The Commonwealth of the Northern Mariana Islands is a group of islands in the Pacific Ocean. Its capital is Saipan.",
"Washington, D.C. is the capital of the United States.",
"Capital punishment has existed in the United States since before it was a country."
],
"top_n": 3
}'
```
</TabItem>
</Tabs>
See [Bedrock Rerank](./bedrock_rerank) for using Bedrock's Rerank API in the Cohere `/rerank` format.
## Bedrock Application Inference Profile
@ -2490,38 +2300,6 @@ model_list:
</Tabs>
Text to Image :
```bash
curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \
-d '{
"model": "amazon.nova-canvas-v1:0",
"prompt": "A cute baby sea otter"
}'
```
Color Guided Generation:
```bash
curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \
-d '{
"model": "amazon.nova-canvas-v1:0",
"prompt": "A cute baby sea otter",
"taskType": "COLOR_GUIDED_GENERATION",
"colorGuidedGenerationParams":{"colors":["#FFFFFF"]}
}'
```
| Model Name | Function Call |
|-------------------------|---------------------------------------------|
| Stable Diffusion 3 - v0 | `image_generation(model="bedrock/stability.stability.sd3-large-v1:0", prompt=prompt)` |
| Stable Diffusion - v0 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v0", prompt=prompt)` |
| Stable Diffusion - v1 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v1", prompt=prompt)` |
| Amazon Nova Canvas - v0 | `image_generation(model="bedrock/amazon.nova-canvas-v1:0", prompt=prompt)` |
### Passing an external BedrockRuntime.Client as a parameter - Completion()
This is a deprecated flow. Boto3 is not async. And boto3.client does not let us make the http call through httpx. Pass in your aws params through the method above 👆. [See Auth Code](https://github.com/BerriAI/litellm/blob/55a20c7cce99a93d36a82bf3ae90ba3baf9a7f89/litellm/llms/bedrock_httpx.py#L284) [Add new auth flow](https://github.com/BerriAI/litellm/issues)

View file

@ -9,6 +9,7 @@ Use Amazon Bedrock Batch Inference API through LiteLLM.
|----------|---------|
| Description | Amazon Bedrock Batch Inference allows you to run inference on large datasets asynchronously |
| Provider Doc | [AWS Bedrock Batch Inference ↗](https://docs.aws.amazon.com/bedrock/latest/userguide/batch-inference.html) |
| Cost Tracking | ✅ Supported |
## Overview

View file

@ -2,11 +2,11 @@
## Supported Embedding Models
| Provider | LiteLLM Route | AWS Documentation |
|----------|---------------|-------------------|
| Amazon Titan | `bedrock/amazon.*` | [Amazon Titan Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/titan-embedding-models.html) |
| Cohere | `bedrock/cohere.*` | [Cohere Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-embed.html) |
| TwelveLabs | `bedrock/us.twelvelabs.*` | [TwelveLabs](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-twelvelabs.html) |
| Provider | LiteLLM Route | AWS Documentation | Cost Tracking |
|----------|---------------|-------------------|---------------|
| Amazon Titan | `bedrock/amazon.*` | [Amazon Titan Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/titan-embedding-models.html) | ✅ |
| Cohere | `bedrock/cohere.*` | [Cohere Embeddings](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-cohere-embed.html) | ✅ |
| TwelveLabs | `bedrock/us.twelvelabs.*` | [TwelveLabs](https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-twelvelabs.html) | ✅ |
## Async Invoke Support

View file

@ -0,0 +1,150 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# AWS Bedrock - Image Generation
Use Bedrock for image generation with Stable Diffusion, Amazon Titan Image Generator, and Amazon Nova Canvas models.
## Supported Models
| Model Name | Function Call | Cost Tracking |
|-------------------------|---------------------------------------------|---------------|
| Stable Diffusion 3 - v0 | `image_generation(model="bedrock/stability.stability.sd3-large-v1:0", prompt=prompt)` | ✅ |
| Stable Diffusion - v0 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v0", prompt=prompt)` | ✅ |
| Stable Diffusion - v1 | `image_generation(model="bedrock/stability.stable-diffusion-xl-v1", prompt=prompt)` | ✅ |
| Amazon Titan Image Generator - v1 | `image_generation(model="bedrock/amazon.titan-image-generator-v1", prompt=prompt)` | ✅ |
| Amazon Titan Image Generator - v2 | `image_generation(model="bedrock/amazon.titan-image-generator-v2:0", prompt=prompt)` | ✅ |
| Amazon Nova Canvas - v1 | `image_generation(model="bedrock/amazon.nova-canvas-v1:0", prompt=prompt)` | ✅ |
## Usage
<Tabs>
<TabItem value="sdk" label="SDK">
### Basic Usage
```python
import os
from litellm import image_generation
os.environ["AWS_ACCESS_KEY_ID"] = ""
os.environ["AWS_SECRET_ACCESS_KEY"] = ""
os.environ["AWS_REGION_NAME"] = ""
response = image_generation(
prompt="A cute baby sea otter",
model="bedrock/stability.stable-diffusion-xl-v0",
)
print(f"response: {response}")
```
### Set Optional Parameters
```python
import os
from litellm import image_generation
os.environ["AWS_ACCESS_KEY_ID"] = ""
os.environ["AWS_SECRET_ACCESS_KEY"] = ""
os.environ["AWS_REGION_NAME"] = ""
response = image_generation(
prompt="A cute baby sea otter",
model="bedrock/stability.stable-diffusion-xl-v0",
### OPENAI-COMPATIBLE ###
size="128x512", # width=128, height=512
### PROVIDER-SPECIFIC ### see `AmazonStabilityConfig` in bedrock.py for all params
seed=30
)
print(f"response: {response}")
```
</TabItem>
<TabItem value="proxy" label="PROXY">
### 1. Setup config.yaml
```yaml
model_list:
- model_name: amazon.nova-canvas-v1:0
litellm_params:
model: bedrock/amazon.nova-canvas-v1:0
aws_region_name: "us-east-1"
aws_secret_access_key: my-key # OPTIONAL - all boto3 auth params supported
aws_secret_access_id: my-id # OPTIONAL - all boto3 auth params supported
```
### 2. Start proxy
```bash
litellm --config /path/to/config.yaml
```
### 3. Test it!
**Text to Image:**
```bash
curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \
-d '{
"model": "amazon.nova-canvas-v1:0",
"prompt": "A cute baby sea otter"
}'
```
**Color Guided Generation:**
```bash
curl -L -X POST 'http://0.0.0.0:4000/v1/images/generations' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer $LITELLM_VIRTUAL_KEY' \
-d '{
"model": "amazon.nova-canvas-v1:0",
"prompt": "A cute baby sea otter",
"taskType": "COLOR_GUIDED_GENERATION",
"colorGuidedGenerationParams":{"colors":["#FFFFFF"]}
}'
```
</TabItem>
</Tabs>
## Using Inference Profiles with Image Generation
For AWS Bedrock Application Inference Profiles with image generation, use the `model_id` parameter to specify the inference profile ARN:
<Tabs>
<TabItem value="sdk" label="SDK">
```python
from litellm import image_generation
response = image_generation(
model="bedrock/amazon.nova-canvas-v1:0",
model_id="arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0",
prompt="A cute baby sea otter"
)
print(f"response: {response}")
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```yaml
model_list:
- model_name: nova-canvas-inference-profile
litellm_params:
model: bedrock/amazon.nova-canvas-v1:0
model_id: arn:aws:bedrock:eu-west-1:000000000000:application-inference-profile/a0a0a0a0a0a0
aws_region_name: "eu-west-1"
```
</TabItem>
</Tabs>
## Authentication
All standard Bedrock authentication methods are supported for image generation. See [Bedrock Authentication](./bedrock#boto3---authentication) for details.

View file

@ -0,0 +1,94 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# AWS Bedrock - Rerank API
Use Bedrock's Rerank API in the Cohere `/rerank` format.
:::info Cost Tracking
✅ **Cost tracking is supported** for Bedrock Rerank API calls.
:::
## Supported Parameters
- `model` - the foundation model ARN
- `query` - the query to rerank against
- `documents` - the list of documents to rerank
- `top_n` - the number of results to return
## Usage
<Tabs>
<TabItem label="SDK" value="sdk">
```python
from litellm import rerank
import os
os.environ["AWS_ACCESS_KEY_ID"] = ""
os.environ["AWS_SECRET_ACCESS_KEY"] = ""
os.environ["AWS_REGION_NAME"] = ""
response = rerank(
model="bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0", # provide the model ARN - get this here https://boto3.amazonaws.com/v1/documentation/api/latest/reference/services/bedrock/client/list_foundation_models.html
query="hello",
documents=["hello", "world"],
top_n=2,
)
print(response)
```
</TabItem>
<TabItem label="PROXY" value="proxy">
### 1. Setup config.yaml
```yaml
model_list:
- model_name: bedrock-rerank
litellm_params:
model: bedrock/arn:aws:bedrock:us-west-2::foundation-model/amazon.rerank-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: os.environ/AWS_REGION_NAME
```
### 2. Start proxy server
```bash
litellm --config config.yaml
# RUNNING on http://0.0.0.0:4000
```
### 3. Test it!
```bash
curl http://0.0.0.0:4000/rerank \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"model": "bedrock-rerank",
"query": "What is the capital of the United States?",
"documents": [
"Carson City is the capital city of the American state of Nevada.",
"The Commonwealth of the Northern Mariana Islands is a group of islands in the Pacific Ocean. Its capital is Saipan.",
"Washington, D.C. is the capital of the United States.",
"Capital punishment has existed in the United States since before it was a country."
],
"top_n": 3
}'
```
</TabItem>
</Tabs>
## Authentication
All standard Bedrock authentication methods are supported for rerank. See [Bedrock Authentication](./bedrock#boto3---authentication) for details.

View file

@ -484,6 +484,8 @@ const sidebars = {
items: [
"providers/bedrock",
"providers/bedrock_embedding",
"providers/bedrock_image_gen",
"providers/bedrock_rerank",
"providers/bedrock_agents",
"providers/bedrock_batches",
"providers/bedrock_vector_store",

View file

@ -227,7 +227,9 @@ class OpenTelemetry(CustomLogger):
PeriodicExportingMetricReader,
)
normalized_endpoint = self._normalize_otel_endpoint(self.config.endpoint, 'metrics')
normalized_endpoint = self._normalize_otel_endpoint(
self.config.endpoint, "metrics"
)
_metric_exporter = OTLPMetricExporter(
endpoint=normalized_endpoint,
headers=OpenTelemetry._get_headers_dictionary(self.config.headers),
@ -664,7 +666,9 @@ class OpenTelemetry(CustomLogger):
# Get the resource from the logger provider
logger_provider = get_logger_provider()
resource = getattr(logger_provider, '_resource', None) or _get_litellm_resource()
resource = (
getattr(logger_provider, "_resource", None) or _get_litellm_resource()
)
parent_ctx = span.get_span_context()
provider = (kwargs.get("litellm_params") or {}).get(
@ -1302,7 +1306,9 @@ class OpenTelemetry(CustomLogger):
"OpenTelemetry: intiializing http exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
)
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, 'traces')
normalized_endpoint = self._normalize_otel_endpoint(
self.OTEL_ENDPOINT, "traces"
)
return BatchSpanProcessor(
OTLPSpanExporterHTTP(
endpoint=normalized_endpoint, headers=_split_otel_headers
@ -1313,7 +1319,9 @@ class OpenTelemetry(CustomLogger):
"OpenTelemetry: intiializing grpc exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
)
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, 'traces')
normalized_endpoint = self._normalize_otel_endpoint(
self.OTEL_ENDPOINT, "traces"
)
return BatchSpanProcessor(
OTLPSpanExporterGRPC(
endpoint=normalized_endpoint, headers=_split_otel_headers
@ -1340,7 +1348,7 @@ class OpenTelemetry(CustomLogger):
_split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)
# Normalize endpoint for logs - ensure it points to /v1/logs instead of /v1/traces
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, 'logs')
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "logs")
verbose_logger.debug(
"OpenTelemetry: Log endpoint normalized from %s to %s",
@ -1358,6 +1366,7 @@ class OpenTelemetry(CustomLogger):
if self.OTEL_EXPORTER == "console":
from opentelemetry.sdk._logs.export import ConsoleLogExporter
verbose_logger.debug(
"OpenTelemetry: Using console log exporter. Value of OTEL_EXPORTER: %s",
self.OTEL_EXPORTER,
@ -1368,7 +1377,10 @@ class OpenTelemetry(CustomLogger):
or self.OTEL_EXPORTER == "http/protobuf"
or self.OTEL_EXPORTER == "http/json"
):
from opentelemetry.exporter.otlp.proto.http._log_exporter import OTLPLogExporter
from opentelemetry.exporter.otlp.proto.http._log_exporter import (
OTLPLogExporter,
)
verbose_logger.debug(
"OpenTelemetry: Using HTTP log exporter. Value of OTEL_EXPORTER: %s, endpoint: %s",
self.OTEL_EXPORTER,
@ -1378,7 +1390,10 @@ class OpenTelemetry(CustomLogger):
endpoint=normalized_endpoint, headers=_split_otel_headers
)
elif self.OTEL_EXPORTER == "otlp_grpc" or self.OTEL_EXPORTER == "grpc":
from opentelemetry.exporter.otlp.proto.grpc._log_exporter import OTLPLogExporter
from opentelemetry.exporter.otlp.proto.grpc._log_exporter import (
OTLPLogExporter,
)
verbose_logger.debug(
"OpenTelemetry: Using gRPC log exporter. Value of OTEL_EXPORTER: %s, endpoint: %s",
self.OTEL_EXPORTER,
@ -1393,12 +1408,11 @@ class OpenTelemetry(CustomLogger):
self.OTEL_EXPORTER,
)
from opentelemetry.sdk._logs.export import ConsoleLogExporter
return ConsoleLogExporter()
def _normalize_otel_endpoint(
self,
endpoint: Optional[str],
signal_type: str
self, endpoint: Optional[str], signal_type: str
) -> Optional[str]:
"""
Normalize the endpoint URL for a specific OpenTelemetry signal type.
@ -1431,37 +1445,37 @@ class OpenTelemetry(CustomLogger):
return endpoint
# Validate signal_type
valid_signals = {'traces', 'metrics', 'logs'}
valid_signals = {"traces", "metrics", "logs"}
if signal_type not in valid_signals:
verbose_logger.warning(
"Invalid signal_type '%s' provided to _normalize_otel_endpoint. "
"Valid values: %s. Returning endpoint unchanged.",
signal_type,
valid_signals
valid_signals,
)
return endpoint
# Remove trailing slash
endpoint = endpoint.rstrip('/')
endpoint = endpoint.rstrip("/")
# Check if endpoint already ends with the correct signal path
target_path = f'/v1/{signal_type}'
target_path = f"/v1/{signal_type}"
if endpoint.endswith(target_path):
return endpoint
# Replace existing signal path with the target signal path
other_signals = valid_signals - {signal_type}
for other_signal in other_signals:
other_path = f'/v1/{other_signal}'
other_path = f"/v1/{other_signal}"
if endpoint.endswith(other_path):
endpoint = endpoint.rsplit('/', 1)[0] + f'/{signal_type}'
endpoint = endpoint.rsplit("/", 1)[0] + f"/{signal_type}"
return endpoint
# No existing signal path found, append the target path
if not endpoint.endswith('/v1'):
if not endpoint.endswith("/v1"):
endpoint = endpoint + target_path
else:
endpoint = endpoint + f'/{signal_type}'
endpoint = endpoint + f"/{signal_type}"
return endpoint
@ -1475,11 +1489,10 @@ class OpenTelemetry(CustomLogger):
if isinstance(headers, str):
# when passed HEADERS="x-honeycomb-team=B85YgLm96******"
# Split only on first '=' occurrence
parts = headers.split("=", 1)
if len(parts) == 2:
_split_otel_headers = {parts[0]: parts[1]}
else:
_split_otel_headers = {}
parts = headers.split(",")
for part in parts:
key, value = part.split("=", 1)
_split_otel_headers[key] = value
elif isinstance(headers, dict):
_split_otel_headers = headers
return _split_otel_headers

View file

@ -0,0 +1,160 @@
"""
Transformation logic for Amazon Titan Image Generation.
"""
import types
from typing import List, Optional
from openai.types.image import Image
from litellm import get_model_info
from litellm.types.llms.bedrock import (
AmazonNovaCanvasImageGenerationConfig,
AmazonTitanImageGenerationRequestBody,
AmazonTitanTextToImageParams,
)
from litellm.types.utils import ImageResponse
class AmazonTitanImageGenerationConfig:
"""
Reference: https://us-west-2.console.aws.amazon.com/bedrock/home?region=us-west-2#/providers?model=stability.stable-diffusion-xl-v0
"""
cfg_scale: Optional[int] = None
seed: Optional[float] = None
steps: Optional[List[str]] = None
width: Optional[int] = None
height: Optional[int] = None
def __init__(
self,
cfg_scale: Optional[int] = None,
seed: Optional[float] = None,
steps: Optional[List[str]] = None,
width: Optional[int] = None,
height: Optional[int] = None,
) -> None:
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
setattr(self.__class__, key, value)
@classmethod
def get_config(cls):
return {
k: v
for k, v in cls.__dict__.items()
if not k.startswith("__")
and not isinstance(
v,
(
types.FunctionType,
types.BuiltinFunctionType,
classmethod,
staticmethod,
),
)
and v is not None
}
@classmethod
def _is_titan_model(cls, model: Optional[str] = None) -> bool:
"""
Returns True if the model is a Titan model
Titan models follow this pattern:
"""
if model and "amazon.titan" in model:
return True
return False
@classmethod
def get_supported_openai_params(cls, model: Optional[str] = None) -> List:
return ["size", "n", "quality"]
@classmethod
def map_openai_params(
cls,
non_default_params: dict,
optional_params: dict,
):
from typing import Any, Dict
image_generation_config: Dict[str, Any] = {}
for k, v in non_default_params.items():
if k == "size" and v is not None:
width, height = v.split("x")
image_generation_config["width"] = int(width)
image_generation_config["height"] = int(height)
elif k == "n" and v is not None:
image_generation_config["numberOfImages"] = v
elif (
k == "quality" and v is not None
): # 'auto', 'hd', 'standard', 'high', 'medium', 'low'
if v in ("hd", "premium", "high"):
image_generation_config["quality"] = "premium"
elif v in ("standard", "medium", "low"):
image_generation_config["quality"] = "standard"
if image_generation_config:
optional_params["imageGenerationConfig"] = image_generation_config
return optional_params
@classmethod
def _transform_request(
cls,
input: str,
optional_params: dict,
) -> AmazonTitanImageGenerationRequestBody:
from typing import Any, Dict
image_generation_config = optional_params.pop("imageGenerationConfig", {})
negative_text = optional_params.pop("negativeText", None)
text_to_image_params: Dict[str, Any] = {"text": input}
if negative_text:
text_to_image_params["negativeText"] = negative_text
task_type = optional_params.pop("taskType", "TEXT_IMAGE")
user_specified_image_generation_config = optional_params.pop(
"imageGenerationConfig", {}
)
image_generation_config = {
**image_generation_config,
**user_specified_image_generation_config,
}
return AmazonTitanImageGenerationRequestBody(
taskType=task_type,
textToImageParams=AmazonTitanTextToImageParams(**text_to_image_params), # type: ignore
imageGenerationConfig=AmazonNovaCanvasImageGenerationConfig(
**image_generation_config
),
)
@classmethod
def transform_response_dict_to_openai_response(
cls, model_response: ImageResponse, response_dict: dict
) -> ImageResponse:
image_list: List[Image] = []
for image in response_dict["images"]:
_image = Image(b64_json=image)
image_list.append(_image)
model_response.data = image_list
return model_response
@classmethod
def cost_calculator(
cls,
model: str,
image_response: ImageResponse,
size: Optional[str] = None,
optional_params: Optional[dict] = None,
) -> float:
model_info = get_model_info(model=model)
output_cost_per_image = model_info.get("output_cost_per_image") or 0.0
if not image_response.data:
return 0.0
num_images = len(image_response.data)
return output_cost_per_image * num_images

View file

@ -1,6 +1,9 @@
from typing import Optional
import litellm
from litellm.llms.bedrock.image.amazon_titan_transformation import (
AmazonTitanImageGenerationConfig,
)
from litellm.types.utils import ImageResponse
@ -17,6 +20,13 @@ def cost_calculator(
"""
if litellm.AmazonStability3Config()._is_stability_3_model(model=model):
pass
elif AmazonTitanImageGenerationConfig._is_titan_model(model=model):
return AmazonTitanImageGenerationConfig.cost_calculator(
model=model,
image_response=image_response,
size=size,
optional_params=optional_params,
)
else:
# Stability 1 models
optional_params = optional_params or {}

View file

@ -9,6 +9,15 @@ from pydantic import BaseModel
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
from litellm.llms.bedrock.image.amazon_nova_canvas_transformation import (
AmazonNovaCanvasConfig,
)
from litellm.llms.bedrock.image.amazon_stability3_transformation import (
AmazonStability3Config,
)
from litellm.llms.bedrock.image.amazon_titan_transformation import (
AmazonTitanImageGenerationConfig,
)
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
@ -63,7 +72,7 @@ class BedrockImageGeneration(BaseAWSLLM):
extra_headers=extra_headers,
logging_obj=logging_obj,
prompt=prompt,
api_key=api_key
api_key=api_key,
)
if aimg_generation is True:
@ -190,7 +199,7 @@ class BedrockImageGeneration(BaseAWSLLM):
body = json.dumps(data).encode("utf-8")
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=boto3_credentials_info.credentials,
@ -201,7 +210,7 @@ class BedrockImageGeneration(BaseAWSLLM):
headers=headers,
api_key=api_key,
)
## LOGGING
logging_obj.pre_call(
input=prompt,
@ -306,15 +315,21 @@ class BedrockImageGeneration(BaseAWSLLM):
if response_dict is None:
raise ValueError("Error in response object format, got None")
config_class = (
litellm.AmazonStability3Config
if litellm.AmazonStability3Config._is_stability_3_model(model=model)
else (
litellm.AmazonNovaCanvasConfig
if litellm.AmazonNovaCanvasConfig._is_nova_model(model=model)
else litellm.AmazonStabilityConfig
)
)
config_class: Union[
type[AmazonTitanImageGenerationConfig],
type[AmazonNovaCanvasConfig],
type[AmazonStability3Config],
type[litellm.AmazonStabilityConfig],
]
if AmazonTitanImageGenerationConfig._is_titan_model(model=model):
config_class = AmazonTitanImageGenerationConfig
elif AmazonNovaCanvasConfig._is_nova_model(model=model):
config_class = AmazonNovaCanvasConfig
elif AmazonStability3Config._is_stability_3_model(model=model):
config_class = AmazonStability3Config
else:
config_class = litellm.AmazonStabilityConfig
config_class.transform_response_dict_to_openai_response(
model_response=model_response,
response_dict=response_dict,

View file

@ -336,6 +336,24 @@
"output_cost_per_token": 0.0,
"output_vector_size": 1024
},
"amazon.titan-image-generator-v1": {
"input_cost_per_image": 0.0,
"output_cost_per_image": 0.008,
"output_cost_per_image_premium_image": 0.01,
"output_cost_per_image_above_512_and_512_pixels": 0.01,
"output_cost_per_image_above_512_and_512_pixels_and_premium_image": 0.012,
"litellm_provider": "bedrock",
"mode": "image_generation"
},
"amazon.titan-image-generator-v2": {
"input_cost_per_image": 0.0,
"output_cost_per_image": 0.008,
"output_cost_per_image_premium_image": 0.01,
"output_cost_per_image_above_1024_and_1024_pixels": 0.01,
"output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": 0.012,
"litellm_provider": "bedrock",
"mode": "image_generation"
},
"twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"litellm_provider": "bedrock",

View file

@ -1,12 +1,7 @@
import json
from typing import Any, Dict, List, Literal, Optional, Union
from typing_extensions import (
TYPE_CHECKING,
Required,
TypedDict,
override,
)
from typing_extensions import TYPE_CHECKING, Required, TypedDict, override
from .openai import ChatCompletionToolCallChunk
@ -468,6 +463,15 @@ class AmazonStability3TextToImageResponse(TypedDict, total=False):
finish_reasons: List[str]
class AmazonTitanTextToImageParams(TypedDict, total=False):
"""
Params for Amazon Titan Text to Image API
"""
text: Required[str]
negativeText: str
class AmazonNovaCanvasRequestBase(TypedDict, total=False):
"""
Base class for Amazon Nova Canvas API requests
@ -577,6 +581,16 @@ class AmazonNovaCanvasInpaintingRequest(
imageGenerationConfig: AmazonNovaCanvasImageGenerationConfig
class AmazonTitanImageGenerationRequestBody(TypedDict, total=False):
"""
Config for Amazon Titan Image Generation API
"""
taskType: Literal["TEXT_IMAGE", "COLOR_GUIDED_GENERATION", "INPAINTING"]
textToImageParams: AmazonTitanTextToImageParams
imageGenerationConfig: AmazonNovaCanvasImageGenerationConfig
if TYPE_CHECKING:
from botocore.awsrequest import AWSPreparedRequest
else:

View file

@ -336,6 +336,24 @@
"output_cost_per_token": 0.0,
"output_vector_size": 1024
},
"amazon.titan-image-generator-v1": {
"input_cost_per_image": 0.0,
"output_cost_per_image": 0.008,
"output_cost_per_image_premium_image": 0.01,
"output_cost_per_image_above_512_and_512_pixels": 0.01,
"output_cost_per_image_above_512_and_512_pixels_and_premium_image": 0.012,
"litellm_provider": "bedrock",
"mode": "image_generation"
},
"amazon.titan-image-generator-v2": {
"input_cost_per_image": 0.0,
"output_cost_per_image": 0.008,
"output_cost_per_image_premium_image": 0.01,
"output_cost_per_image_above_1024_and_1024_pixels": 0.01,
"output_cost_per_image_above_1024_and_1024_pixels_and_premium_image": 0.012,
"litellm_provider": "bedrock",
"mode": "image_generation"
},
"twelvelabs.marengo-embed-2-7-v1:0": {
"input_cost_per_token": 7e-05,
"litellm_provider": "bedrock",

View file

@ -508,3 +508,19 @@ def test_backward_compatibility_regular_nova_model():
assert result["taskType"] == "TEXT_IMAGE"
assert result["textToImageParams"]["text"] == prompt
assert result["imageGenerationConfig"]["cfg_scale"] == 7
def test_amazon_titan_image_gen():
from litellm import image_generation
model_id = "bedrock/amazon.titan-image-generator-v1"
response = litellm.image_generation(
model=model_id,
prompt="A serene mountain landscape at sunset with a lake reflection",
aws_region_name="us-east-1",
)
print(f"response cost: {response._hidden_params['response_cost']}")
assert response._hidden_params["response_cost"] > 0

View file

@ -14,26 +14,41 @@ sys.path.insert(
import litellm
from litellm.llms.custom_httpx.http_handler import HTTPHandler, AsyncHTTPHandler
titan_embedding_response = {
"embedding": [0.1, 0.2, 0.3],
"inputTextTokenCount": 10
}
titan_embedding_response = {"embedding": [0.1, 0.2, 0.3], "inputTextTokenCount": 10}
cohere_embedding_response = {
"embeddings": [[0.1, 0.2, 0.3]],
"inputTextTokenCount": 10
}
cohere_embedding_response = {"embeddings": [[0.1, 0.2, 0.3]], "inputTextTokenCount": 10}
img_base_64 = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAGQAAABkBAMAAACCzIhnAAAAG1BMVEURAAD///+ln5/h39/Dv79qX18uHx+If39MPz9oMSdmAAAACXBIWXMAAA7EAAAOxAGVKw4bAAABB0lEQVRYhe2SzWrEIBCAh2A0jxEs4j6GLDS9hqWmV5Flt0cJS+lRwv742DXpEjY1kOZW6HwHFZnPmVEBEARBEARB/jd0KYA/bcUYbPrRLh6amXHJ/K+ypMoyUaGthILzw0l+xI0jsO7ZcmCcm4ILd+QuVYgpHOmDmz6jBeJImdcUCmeBqQpuqRIbVmQsLCrAalrGpfoEqEogqbLTWuXCPCo+Ki1XGqgQ+jVVuhB8bOaHkvmYuzm/b0KYLWwoK58oFqi6XfxQ4Uz7d6WeKpna6ytUs5e8betMcqAv5YPC5EZB2Lm9FIn0/VP6R58+/GEY1X1egVoZ/3bt/EqF6malgSAIgiDIH+QL41409QMY0LMAAAAASUVORK5CYII="
@pytest.mark.parametrize(
"model,input_type,embed_response",
[
("bedrock/amazon.titan-embed-text-v1", "text", titan_embedding_response), # V1 text model
("bedrock/amazon.titan-embed-text-v2:0", "text", titan_embedding_response), # V2 text model
("bedrock/amazon.titan-embed-image-v1", "image", titan_embedding_response), # Image model
("bedrock/cohere.embed-english-v3", "text", cohere_embedding_response), # Cohere English
("bedrock/cohere.embed-multilingual-v3", "text", cohere_embedding_response), # Cohere Multilingual
(
"bedrock/amazon.titan-embed-text-v1",
"text",
titan_embedding_response,
), # V1 text model
(
"bedrock/amazon.titan-embed-text-v2:0",
"text",
titan_embedding_response,
), # V2 text model
(
"bedrock/amazon.titan-embed-image-v1",
"image",
titan_embedding_response,
), # Image model
(
"bedrock/cohere.embed-english-v3",
"text",
cohere_embedding_response,
), # Cohere English
(
"bedrock/cohere.embed-multilingual-v3",
"text",
cohere_embedding_response,
), # Cohere Multilingual
],
)
def test_bedrock_embedding_models(model, input_type, embed_response):
@ -49,7 +64,9 @@ def test_bedrock_embedding_models(model, input_type, embed_response):
mock_post.return_value = mock_response
# Prepare input based on type
input_data = img_base_64 if input_type == "image" else "Hello world from litellm"
input_data = (
img_base_64 if input_type == "image" else "Hello world from litellm"
)
try:
response = litellm.embedding(
@ -63,8 +80,8 @@ def test_bedrock_embedding_models(model, input_type, embed_response):
# Verify response structure
assert isinstance(response, litellm.EmbeddingResponse)
print(response.data)
assert isinstance(response.data[0]['embedding'], list)
assert len(response.data[0]['embedding']) == 3 # Based on mock response
assert isinstance(response.data[0]["embedding"], list)
assert len(response.data[0]["embedding"]) == 3 # Based on mock response
# Fetch request body
request_data = json.loads(mock_post.call_args.kwargs["data"])
@ -72,7 +89,9 @@ def test_bedrock_embedding_models(model, input_type, embed_response):
# Verify AWS params are not in request body
aws_params = ["aws_region_name", "aws_bedrock_runtime_endpoint"]
for param in aws_params:
assert param not in request_data, f"AWS param {param} should not be in request body"
assert (
param not in request_data
), f"AWS param {param} should not be in request body"
except Exception as e:
pytest.fail(f"Error occurred: {e}")
@ -85,7 +104,7 @@ def test_e2e_bedrock_embedding():
"""
print("Testing text embedding...")
original_region_name = os.environ.get("AWS_REGION_NAME")
os.environ["AWS_REGION_NAME"] = "us-east-1"
litellm._turn_on_debug()
@ -93,34 +112,47 @@ def test_e2e_bedrock_embedding():
model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0",
input=["Hello world from LiteLLM with TwelveLabs Marengo!"],
)
# Validate response structure
assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type"
assert hasattr(response, 'data'), "Response should have 'data' attribute"
assert isinstance(
response, litellm.EmbeddingResponse
), "Response should be EmbeddingResponse type"
assert hasattr(response, "data"), "Response should have 'data' attribute"
assert len(response.data) > 0, "Response data should not be empty"
# Validate first embedding
embedding_obj = response.data[0]
assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute"
assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats"
assert hasattr(
embedding_obj, "embedding"
), "Embedding object should have 'embedding' attribute"
assert isinstance(
embedding_obj.embedding, list
), "Embedding should be a list of floats"
assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty"
assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric"
assert all(
isinstance(x, (int, float)) for x in embedding_obj.embedding
), "All embedding values should be numeric"
# Validate embedding properties
assert embedding_obj.index == 0, "First embedding should have index 0"
assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'"
assert (
embedding_obj.object == "embedding"
), "Embedding object type should be 'embedding'"
# Validate usage information
assert hasattr(response, 'usage'), "Response should have usage information"
assert hasattr(response, "usage"), "Response should have usage information"
assert response.usage is not None, "Usage should not be None"
assert response.usage.total_tokens >= 0, "Total tokens should be non-negative"
print(f"Text embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}")
print(
f"Text 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_embedding_image_twelvelabs_marengo():
"""
Test image embedding with TwelveLabs Marengo.
@ -130,46 +162,60 @@ def test_e2e_bedrock_embedding_image_twelvelabs_marengo():
original_region_name = os.environ.get("AWS_REGION_NAME")
os.environ["AWS_REGION_NAME"] = "us-east-1"
litellm._turn_on_debug()
# Load duck.png and convert to base64
duck_img_path = os.path.join(os.path.dirname(__file__), "duck.png")
with open(duck_img_path, "rb") as img_file:
duck_img_data = base64.b64encode(img_file.read()).decode('utf-8')
duck_img_data = base64.b64encode(img_file.read()).decode("utf-8")
duck_img_base64 = f"data:image/png;base64,{duck_img_data}"
response = litellm.embedding(
model="bedrock/us.twelvelabs.marengo-embed-2-7-v1:0",
input=[duck_img_base64],
aws_region_name="us-east-1",
input_type="image"
input_type="image",
)
# Validate response structure
assert isinstance(response, litellm.EmbeddingResponse), "Response should be EmbeddingResponse type"
assert hasattr(response, 'data'), "Response should have 'data' attribute"
assert isinstance(
response, litellm.EmbeddingResponse
), "Response should be EmbeddingResponse type"
assert hasattr(response, "data"), "Response should have 'data' attribute"
assert len(response.data) > 0, "Response data should not be empty"
# Validate first embedding
embedding_obj = response.data[0]
assert hasattr(embedding_obj, 'embedding'), "Embedding object should have 'embedding' attribute"
assert isinstance(embedding_obj.embedding, list), "Embedding should be a list of floats"
assert hasattr(
embedding_obj, "embedding"
), "Embedding object should have 'embedding' attribute"
assert isinstance(
embedding_obj.embedding, list
), "Embedding should be a list of floats"
assert len(embedding_obj.embedding) > 0, "Embedding vector should not be empty"
assert all(isinstance(x, (int, float)) for x in embedding_obj.embedding), "All embedding values should be numeric"
assert all(
isinstance(x, (int, float)) for x in embedding_obj.embedding
), "All embedding values should be numeric"
# Validate embedding properties
assert embedding_obj.index == 0, "First embedding should have index 0"
assert embedding_obj.object == "embedding", "Embedding object type should be 'embedding'"
assert (
embedding_obj.object == "embedding"
), "Embedding object type should be 'embedding'"
# Validate usage information
assert hasattr(response, 'usage'), "Response should have usage information"
assert hasattr(response, "usage"), "Response should have usage information"
assert response.usage is not None, "Usage should not be None"
assert response.usage.total_tokens >= 0, "Total tokens should be non-negative"
# TwelveLabs Marengo should return 1024-dimensional embeddings
expected_dimension = 1024
assert len(embedding_obj.embedding) == expected_dimension, f"TwelveLabs Marengo should return {expected_dimension}-dimensional embeddings, got {len(embedding_obj.embedding)}"
print(f"Image embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}")
assert (
len(embedding_obj.embedding) == expected_dimension
), f"TwelveLabs Marengo should return {expected_dimension}-dimensional embeddings, got {len(embedding_obj.embedding)}"
print(
f"Image embedding successful! Vector size: {len(embedding_obj.embedding)}, Response: {response}"
)
# Restore original region name
if original_region_name:
@ -185,37 +231,54 @@ def test_e2e_bedrock_async_invoke_embedding_twelvelabs_marengo():
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:
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/"
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 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"
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}")
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
@ -231,32 +294,45 @@ async def test_e2e_bedrock_async_invoke_embedding_async_twelvelabs_marengo():
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:
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/"
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 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-async-job-456", "Invocation ARN should be preserved"
print(f"Async invoke embedding successful! Invocation ARN: {response._hidden_params._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-async-job-456"
), "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
os.environ["AWS_REGION_NAME"] = original_region_name

View file

@ -1,24 +1,24 @@
import json
import os
import sys
import unittest
from unittest.mock import MagicMock, patch
from datetime import datetime, timedelta
import time
import unittest
from datetime import datetime, timedelta
from unittest.mock import MagicMock, patch
# Adds the grandparent directory to sys.path to allow importing project modules
sys.path.insert(0, os.path.abspath("../.."))
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk._logs import LoggerProvider as OTLoggerProvider
from opentelemetry.sdk._logs.export import SimpleLogRecordProcessor, InMemoryLogExporter
from opentelemetry.sdk._logs.export import InMemoryLogExporter, SimpleLogRecordProcessor
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.metrics.export import InMemoryMetricReader
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
class TestOpenTelemetryGuardrails(unittest.TestCase):
@patch("litellm.integrations.opentelemetry.datetime")
@ -604,19 +604,35 @@ class TestOpenTelemetry(unittest.TestCase):
logs = self.wait_for_log(log_exporter, "gen_ai.")
self.assertTrue(logs, "Expected at least one gen_ai log")
user_logs = [log for log in logs if log.log_record.attributes.get("event_name") == "gen_ai.content.prompt"]
user_logs = [
log
for log in logs
if log.log_record.attributes.get("event_name") == "gen_ai.content.prompt"
]
self.assertTrue(user_logs, "did not see a gen_ai.content.prompt log")
# check log bodies
user_prompt = user_logs[0].log_record.attributes.get("gen_ai.prompt")
self.assertEqual("What is the capital of France?", user_prompt, "did not see a prompt message")
self.assertEqual(
"What is the capital of France?",
user_prompt,
"did not see a prompt message",
)
choice_logs = [log for log in logs if log.log_record.attributes.get("event_name") == "gen_ai.content.completion"]
choice_logs = [
log
for log in logs
if log.log_record.attributes.get("event_name")
== "gen_ai.content.completion"
]
self.assertTrue(choice_logs, "did not see a gen_ai.content.completion event")
choice_response = choice_logs[0].log_record.body
self.assertIsNotNone(choice_response, "did not see a response message")
self.assertEqual("stop", choice_response.get("finish_reason"), "did not see expected finish reason")
self.assertEqual(
"stop",
choice_response.get("finish_reason"),
"did not see expected finish reason",
)
def test_handle_success_spans_only(self):
# make sure neither events nor metrics is on
@ -670,8 +686,7 @@ class TestOpenTelemetry(unittest.TestCase):
# )
# model attribute should be on that span
found = any(
s.attributes
and s.attributes.get("gen_ai.request.model") == self.MODEL
s.attributes and s.attributes.get("gen_ai.request.model") == self.MODEL
for s in spans
)
self.assertTrue(found, "expected gen_ai.request.model on span attributes")
@ -755,13 +770,7 @@ class TestOpenTelemetry(unittest.TestCase):
def test_get_span_name_with_generation_name(self):
"""Test _get_span_name returns generation_name when present"""
otel = OpenTelemetry()
kwargs = {
"litellm_params": {
"metadata": {
"generation_name": "custom_span"
}
}
}
kwargs = {"litellm_params": {"metadata": {"generation_name": "custom_span"}}}
result = otel._get_span_name(kwargs)
self.assertEqual(result, "custom_span")
@ -774,7 +783,7 @@ class TestOpenTelemetry(unittest.TestCase):
result = otel._get_span_name(kwargs)
self.assertEqual(result, LITELLM_REQUEST_SPAN_NAME)
@patch('litellm.turn_off_message_logging', False)
@patch("litellm.turn_off_message_logging", False)
def test_maybe_log_raw_request_creates_span(self):
"""Test _maybe_log_raw_request creates span when logging enabled"""
from litellm.integrations.opentelemetry import RAW_REQUEST_SPAN_NAME
@ -790,12 +799,16 @@ class TestOpenTelemetry(unittest.TestCase):
otel._to_ns = MagicMock(return_value=1234567890)
kwargs = {"litellm_params": {"metadata": {}}}
otel._maybe_log_raw_request(kwargs, {}, datetime.now(), datetime.now(), MagicMock())
otel._maybe_log_raw_request(
kwargs, {}, datetime.now(), datetime.now(), MagicMock()
)
mock_tracer.start_span.assert_called_once()
self.assertEqual(mock_tracer.start_span.call_args[1]['name'], RAW_REQUEST_SPAN_NAME)
self.assertEqual(
mock_tracer.start_span.call_args[1]["name"], RAW_REQUEST_SPAN_NAME
)
@patch('litellm.turn_off_message_logging', True)
@patch("litellm.turn_off_message_logging", True)
def test_maybe_log_raw_request_skips_when_logging_disabled(self):
"""Test _maybe_log_raw_request skips when logging disabled"""
otel = OpenTelemetry()
@ -803,24 +816,50 @@ class TestOpenTelemetry(unittest.TestCase):
otel.get_tracer_to_use_for_request = MagicMock(return_value=mock_tracer)
kwargs = {"litellm_params": {"metadata": {}}}
otel._maybe_log_raw_request(kwargs, {}, datetime.now(), datetime.now(), MagicMock())
otel._maybe_log_raw_request(
kwargs, {}, datetime.now(), datetime.now(), MagicMock()
)
mock_tracer.start_span.assert_not_called()
class TestOpenTelemetryHeaderSplitting(unittest.TestCase):
"""Test suite for _get_headers_dictionary method"""
def test_split_multiple_headers_comma_separated(self):
"""Test splitting multiple headers separated by commas"""
otel = OpenTelemetry()
headers = "api-key=key,other-config-value=value"
result = otel._get_headers_dictionary(headers)
self.assertEqual(result, {"api-key": "key", "other-config-value": "value"})
def test_split_headers_with_equals_in_values(self):
"""Test splitting headers where values contain equals signs (split only on first '=')"""
otel = OpenTelemetry()
headers = "api-key=value1=part2,config=setting=enabled"
result = otel._get_headers_dictionary(headers)
self.assertEqual(
result, {"api-key": "value1=part2", "config": "setting=enabled"}
)
class TestOpenTelemetryEndpointNormalization(unittest.TestCase):
"""Test suite for the unified _normalize_otel_endpoint method"""
def test_normalize_traces_endpoint_from_logs_path(self):
"""Test normalizing endpoint with /v1/logs to /v1/traces"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/logs", "traces")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/logs", "traces"
)
self.assertEqual(result, "http://collector:4318/v1/traces")
def test_normalize_traces_endpoint_from_metrics_path(self):
"""Test normalizing endpoint with /v1/metrics to /v1/traces"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/metrics", "traces")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/metrics", "traces"
)
self.assertEqual(result, "http://collector:4318/v1/traces")
def test_normalize_traces_endpoint_from_base_url(self):
@ -838,19 +877,25 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase):
def test_normalize_traces_endpoint_already_correct(self):
"""Test endpoint already ending with /v1/traces remains unchanged"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/traces", "traces")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/traces", "traces"
)
self.assertEqual(result, "http://collector:4318/v1/traces")
def test_normalize_metrics_endpoint_from_traces_path(self):
"""Test normalizing endpoint with /v1/traces to /v1/metrics"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/traces", "metrics")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/traces", "metrics"
)
self.assertEqual(result, "http://collector:4318/v1/metrics")
def test_normalize_metrics_endpoint_from_logs_path(self):
"""Test normalizing endpoint with /v1/logs to /v1/metrics"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/logs", "metrics")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/logs", "metrics"
)
self.assertEqual(result, "http://collector:4318/v1/metrics")
def test_normalize_metrics_endpoint_from_base_url(self):
@ -862,19 +907,25 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase):
def test_normalize_metrics_endpoint_already_correct(self):
"""Test endpoint already ending with /v1/metrics remains unchanged"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/metrics", "metrics")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/metrics", "metrics"
)
self.assertEqual(result, "http://collector:4318/v1/metrics")
def test_normalize_logs_endpoint_from_traces_path(self):
"""Test normalizing endpoint with /v1/traces to /v1/logs"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/traces", "logs")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/traces", "logs"
)
self.assertEqual(result, "http://collector:4318/v1/logs")
def test_normalize_logs_endpoint_from_metrics_path(self):
"""Test normalizing endpoint with /v1/metrics to /v1/logs"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/v1/metrics", "logs")
result = otel._normalize_otel_endpoint(
"http://collector:4318/v1/metrics", "logs"
)
self.assertEqual(result, "http://collector:4318/v1/logs")
def test_normalize_logs_endpoint_from_base_url(self):
@ -912,7 +963,7 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase):
otel = OpenTelemetry()
endpoint = "http://collector:4318/v1/traces"
with patch('litellm._logging.verbose_logger.warning') as mock_warning:
with patch("litellm._logging.verbose_logger.warning") as mock_warning:
result = otel._normalize_otel_endpoint(endpoint, "invalid")
# Should return endpoint unchanged
@ -924,18 +975,24 @@ class TestOpenTelemetryEndpointNormalization(unittest.TestCase):
call_args = mock_warning.call_args[0]
self.assertIn("Invalid signal_type", call_args[0])
self.assertEqual(call_args[1], "invalid") # signal_type parameter
self.assertEqual(call_args[2], {'traces', 'metrics', 'logs'}) # valid_signals parameter
self.assertEqual(
call_args[2], {"traces", "metrics", "logs"}
) # valid_signals parameter
def test_normalize_endpoint_https(self):
"""Test normalization works with https URLs"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("https://collector.example.com:4318", "logs")
result = otel._normalize_otel_endpoint(
"https://collector.example.com:4318", "logs"
)
self.assertEqual(result, "https://collector.example.com:4318/v1/logs")
def test_normalize_endpoint_with_path_prefix(self):
"""Test normalization works with URLs that have path prefixes"""
otel = OpenTelemetry()
result = otel._normalize_otel_endpoint("http://collector:4318/otel/v1/traces", "logs")
result = otel._normalize_otel_endpoint(
"http://collector:4318/otel/v1/traces", "logs"
)
# Should replace the final /traces with /logs
self.assertEqual(result, "http://collector:4318/otel/v1/logs")
@ -984,8 +1041,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
from opentelemetry.sdk.trace.export import BatchSpanProcessor
config = OpenTelemetryConfig(
exporter="otlp_http",
endpoint="http://collector:4318"
exporter="otlp_http", endpoint="http://collector:4318"
)
otel = OpenTelemetry(config=config)
@ -1005,8 +1061,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
from opentelemetry.sdk.trace.export import BatchSpanProcessor
config = OpenTelemetryConfig(
exporter="otlp_grpc",
endpoint="http://collector:4317"
exporter="otlp_grpc", endpoint="http://collector:4317"
)
otel = OpenTelemetry(config=config)
@ -1025,10 +1080,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
)
from opentelemetry.sdk.trace.export import BatchSpanProcessor
config = OpenTelemetryConfig(
exporter="grpc",
endpoint="http://collector:4317"
)
config = OpenTelemetryConfig(exporter="grpc", endpoint="http://collector:4317")
otel = OpenTelemetry(config=config)
processor = otel._get_span_processor()
@ -1047,8 +1099,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
from opentelemetry.sdk.trace.export import BatchSpanProcessor
config = OpenTelemetryConfig(
exporter="http/protobuf",
endpoint="http://collector:4318"
exporter="http/protobuf", endpoint="http://collector:4318"
)
otel = OpenTelemetry(config=config)
@ -1083,9 +1134,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
from opentelemetry.exporter.otlp.proto.http._log_exporter import OTLPLogExporter
config = OpenTelemetryConfig(
exporter="otlp_http",
endpoint="http://collector:4318",
enable_events=True
exporter="otlp_http", endpoint="http://collector:4318", enable_events=True
)
otel = OpenTelemetry(config=config)
@ -1095,16 +1144,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
self.assertIsInstance(exporter, OTLPLogExporter)
# Check that it's from the http module by checking the module name
self.assertIn('http', exporter.__class__.__module__)
self.assertIn("http", exporter.__class__.__module__)
def test_get_log_exporter_uses_grpc_exporter_for_otlp_grpc(self):
"""Test that otlp_grpc protocol uses gRPC OTLPLogExporter"""
from opentelemetry.exporter.otlp.proto.grpc._log_exporter import OTLPLogExporter
config = OpenTelemetryConfig(
exporter="otlp_grpc",
endpoint="http://collector:4317",
enable_events=True
exporter="otlp_grpc", endpoint="http://collector:4317", enable_events=True
)
otel = OpenTelemetry(config=config)
@ -1114,16 +1161,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
self.assertIsInstance(exporter, OTLPLogExporter)
# Check that it's from the grpc module by checking the module name
self.assertIn('grpc', exporter.__class__.__module__)
self.assertIn("grpc", exporter.__class__.__module__)
def test_get_log_exporter_uses_grpc_exporter_for_grpc_alias(self):
"""Test that 'grpc' protocol alias uses gRPC OTLPLogExporter"""
from opentelemetry.exporter.otlp.proto.grpc._log_exporter import OTLPLogExporter
config = OpenTelemetryConfig(
exporter="grpc",
endpoint="http://collector:4317",
enable_events=True
exporter="grpc", endpoint="http://collector:4317", enable_events=True
)
otel = OpenTelemetry(config=config)
@ -1133,16 +1178,13 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
self.assertIsInstance(exporter, OTLPLogExporter)
# Check that it's from the grpc module by checking the module name
self.assertIn('grpc', exporter.__class__.__module__)
self.assertIn("grpc", exporter.__class__.__module__)
def test_get_log_exporter_uses_console_exporter_for_console(self):
"""Test that console protocol uses ConsoleLogExporter"""
from opentelemetry.sdk._logs.export import ConsoleLogExporter
config = OpenTelemetryConfig(
exporter="console",
enable_events=True
)
config = OpenTelemetryConfig(exporter="console", enable_events=True)
otel = OpenTelemetry(config=config)
exporter = otel._get_log_exporter()
@ -1154,13 +1196,10 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
"""Test that unknown protocol defaults to ConsoleLogExporter with warning"""
from opentelemetry.sdk._logs.export import ConsoleLogExporter
config = OpenTelemetryConfig(
exporter="unknown_protocol",
enable_events=True
)
config = OpenTelemetryConfig(exporter="unknown_protocol", enable_events=True)
otel = OpenTelemetry(config=config)
with patch('litellm._logging.verbose_logger.warning') as mock_warning:
with patch("litellm._logging.verbose_logger.warning") as mock_warning:
exporter = otel._get_log_exporter()
# Verify the exporter defaults to console
@ -1172,7 +1211,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
self.assertIn("Unknown log exporter", args[0])
self.assertIn("unknown_protocol", args[1])
@patch.dict(os.environ, {"OTEL_EXPORTER": "otlp_http", "OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4318"}, clear=False)
@patch.dict(
os.environ,
{
"OTEL_EXPORTER": "otlp_http",
"OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4318",
},
clear=False,
)
def test_protocol_selection_from_environment_http(self):
"""Test that protocol selection works correctly from environment variables for HTTP"""
from opentelemetry.exporter.otlp.proto.http.trace_exporter import (
@ -1189,7 +1235,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
self.assertIsInstance(processor, BatchSpanProcessor)
self.assertIsInstance(processor.span_exporter, OTLPSpanExporterHTTP)
@patch.dict(os.environ, {"OTEL_EXPORTER": "otlp_grpc", "OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4317"}, clear=False)
@patch.dict(
os.environ,
{
"OTEL_EXPORTER": "otlp_grpc",
"OTEL_EXPORTER_OTLP_ENDPOINT": "http://collector:4317",
},
clear=False,
)
def test_protocol_selection_from_environment_grpc(self):
"""Test that protocol selection works correctly from environment variables for gRPC"""
from opentelemetry.exporter.otlp.proto.grpc.trace_exporter import (
@ -1209,8 +1262,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
def test_http_exporter_endpoint_normalization_for_traces(self):
"""Test that HTTP trace exporter gets properly normalized endpoint"""
config = OpenTelemetryConfig(
exporter="otlp_http",
endpoint="http://collector:4318"
exporter="otlp_http", endpoint="http://collector:4318"
)
otel = OpenTelemetry(config=config)
@ -1218,14 +1270,13 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
# Verify the endpoint was normalized to include /v1/traces
# Access the private _endpoint attribute if available
if hasattr(processor.span_exporter, '_endpoint'):
if hasattr(processor.span_exporter, "_endpoint"):
self.assertEqual(processor.span_exporter._endpoint, "http://collector:4318/v1/traces") # type: ignore[attr-defined]
def test_grpc_exporter_endpoint_normalization_for_traces(self):
"""Test that gRPC trace exporter gets properly normalized endpoint"""
config = OpenTelemetryConfig(
exporter="otlp_grpc",
endpoint="http://collector:4317"
exporter="otlp_grpc", endpoint="http://collector:4317"
)
otel = OpenTelemetry(config=config)
@ -1233,12 +1284,14 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
# Verify the endpoint was normalized to include /v1/traces
# Note: gRPC exporters strip the http:// prefix, so we check for the normalized path
if hasattr(processor.span_exporter, '_endpoint'):
if hasattr(processor.span_exporter, "_endpoint"):
# gRPC exporter strips http:// prefix
self.assertIn('collector:4317', processor.span_exporter._endpoint) # type: ignore[attr-defined]
self.assertIn("collector:4317", processor.span_exporter._endpoint) # type: ignore[attr-defined]
# The endpoint should have been normalized with /v1/traces before being passed to gRPC exporter
# We verify this by checking the normalization function was called correctly
normalized = otel._normalize_otel_endpoint("http://collector:4317", "traces")
normalized = otel._normalize_otel_endpoint(
"http://collector:4317", "traces"
)
self.assertEqual(normalized, "http://collector:4317/v1/traces")
def test_http_log_exporter_endpoint_normalization_for_logs(self):
@ -1246,7 +1299,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
config = OpenTelemetryConfig(
exporter="otlp_http",
endpoint="http://collector:4318/v1/traces",
enable_events=True
enable_events=True,
)
otel = OpenTelemetry(config=config)
@ -1254,7 +1307,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
# Verify the endpoint was normalized to /v1/logs (not /v1/traces)
# Access the private _endpoint attribute if available
if hasattr(exporter, '_endpoint'):
if hasattr(exporter, "_endpoint"):
self.assertEqual(exporter._endpoint, "http://collector:4318/v1/logs") # type: ignore[attr-defined]
def test_grpc_log_exporter_endpoint_normalization_for_logs(self):
@ -1262,7 +1315,7 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
config = OpenTelemetryConfig(
exporter="otlp_grpc",
endpoint="http://collector:4317/v1/traces",
enable_events=True
enable_events=True,
)
otel = OpenTelemetry(config=config)
@ -1270,10 +1323,12 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase):
# Verify the endpoint was normalized to /v1/logs (not /v1/traces)
# Note: gRPC exporters strip the http:// prefix, so we check for the normalized path
if hasattr(exporter, '_endpoint'):
if hasattr(exporter, "_endpoint"):
# gRPC exporter strips http:// prefix
self.assertIn('collector:4317', exporter._endpoint) # type: ignore[attr-defined]
self.assertIn("collector:4317", exporter._endpoint) # type: ignore[attr-defined]
# The endpoint should have been normalized with /v1/logs before being passed to gRPC exporter
# We verify this by checking the normalization function was called correctly
normalized = otel._normalize_otel_endpoint("http://collector:4317/v1/traces", "logs")
normalized = otel._normalize_otel_endpoint(
"http://collector:4317/v1/traces", "logs"
)
self.assertEqual(normalized, "http://collector:4317/v1/logs")